FazBrowse GitHub Viewer | Trending |
URL:
| Home
Tools: [Download Repo ZIP]   [Original HTTPS Page]

add TensorFlowOpLayer. · Liang813/TensorFlow.NET@a0ec655 · GitHub

Commit a0ec655

Browse files
committed
add TensorFlowOpLayer.
1 parent b79d6bc commit a0ec655

12 files changed

Lines changed: 301 additions & 204 deletions

‎src/TensorFlowNET.Core/Keras/ArgsDefinition/NodeArgs.cs‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -11,7 +11,7 @@ public class NodeArgs
1111
public Layer[] InboundLayers { get; set; }
1212
public int[] NodeIndices { get; set; }
1313
public int[] TensorIndices { get; set; }
14-
public Tensor InputTensors { get; set; }
14+
public Tensors InputTensors { get; set; }
1515
public Tensors Outputs { get; set; }
1616
}
1717
}
Lines changed: 11 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,11 @@
1+
using System;
2+
using System.Collections.Generic;
3+
using System.Text;
4+
5+
namespace Tensorflow.Keras.ArgsDefinition
6+
{
7+
public class TensorFlowOpLayerArgs : LayerArgs
8+
{
9+
public NodeDef NodeDef { get; set; }
10+
}
11+
}

‎src/TensorFlowNET.Core/Keras/Engine/BaseLayerUtils.cs‎

Lines changed: 0 additions & 47 deletions
This file was deleted.

‎src/TensorFlowNET.Core/Keras/Engine/Functional.cs‎

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -4,6 +4,7 @@
44
using System.Security.Cryptography.X509Certificates;
55
using System.Text;
66
using Tensorflow.Keras.ArgsDefinition;
7+
using Tensorflow.Keras.Utils;
78

89
namespace Tensorflow.Keras.Engine
910
{
@@ -50,7 +51,7 @@ void _init_graph_network(Tensors inputs, Tensors outputs)
5051
_autocast = false;
5152

5253
if (outputs.Any(x => x.KerasHistory == null))
53-
BaseLayerUtils.CreateKerasHistoryHelper(outputs);
54+
base_layer_utils.create_keras_history(outputs);
5455

5556
// Build self._output_layers:
5657
foreach (var x in outputs)

‎src/TensorFlowNET.Core/Keras/Engine/KerasHistory.cs‎

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -9,7 +9,7 @@ namespace Tensorflow.Keras.Engine
99
/// </summary>
1010
public class KerasHistory
1111
{
12-
Layer layer;
12+
public Layer layer;
1313
int node_index;
1414
int tensor_index;
1515
public Tensor tensor;
@@ -20,6 +20,7 @@ public KerasHistory(Layer layer, int node_index, int tensor_index, Tensor tensor
2020
this.node_index = node_index;
2121
this.tensor_index = tensor_index;
2222
this.tensor = tensor;
23+
Layer.KerasHistories.Add(this);
2324
Console.WriteLine(tensor.name);
2425
}
2526

Lines changed: 65 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,65 @@
1+
using System;
2+
using System.Collections.Generic;
3+
using System.Text;
4+
using Tensorflow.Keras.Utils;
5+
using static Tensorflow.Binding;
6+
7+
namespace Tensorflow.Keras.Engine
8+
{
9+
public partial class Layer
10+
{
11+
protected virtual IVariableV1 add_weight(string name,
12+
TensorShape shape,
13+
TF_DataType dtype = TF_DataType.TF_FLOAT,
14+
IInitializer initializer = null,
15+
IRegularizer regularizer = null,
16+
VariableSynchronization synchronization = VariableSynchronization.Auto,
17+
VariableAggregation aggregation = VariableAggregation.None,
18+
bool trainable = true,
19+
Func<VariableArgs, IVariableV1> getter = null)
20+
{
21+
// Initialize variable when no initializer provided
22+
if (initializer == null)
23+
{
24+
// If dtype is DT_FLOAT, provide a uniform unit scaling initializer
25+
if (dtype.is_floating())
26+
initializer = tf.glorot_uniform_initializer;
27+
else if (dtype.is_integer())
28+
initializer = tf.zeros_initializer;
29+
else
30+
throw new ValueError($"An initializer for variable {name} of type {dtype.as_base_dtype()} is required for layer {name}");
31+
}
32+
33+
if (synchronization == VariableSynchronization.OnRead)
34+
trainable = false;
35+
36+
var args = new VariableArgs
37+
{
38+
Name = name,
39+
Shape = shape,
40+
DType = dtype,
41+
Getter = getter ?? base_layer_utils.make_variable,
42+
Overwrite = true,
43+
Initializer = initializer,
44+
Synchronization = synchronization,
45+
Aggregation = aggregation,
46+
Trainable = trainable
47+
};
48+
var variable = _add_variable_with_custom_getter(args);
49+
50+
if (regularizer != null)
51+
{
52+
var name_in_scope = variable.Name.Split(':')[0];
53+
_handle_weight_regularization(name_in_scope, variable, regularizer);
54+
}
55+
56+
//backend.track_variable(variable);
57+
if (trainable == true)
58+
trainableWeights.Add(variable);
59+
else
60+
nonTrainableWeights.Add(variable);
61+
62+
return variable;
63+
}
64+
}
65+
}
Lines changed: 62 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,62 @@
1+
using System;
2+
using System.Collections.Generic;
3+
using System.Linq;
4+
using System.Text;
5+
using System.Threading;
6+
using Tensorflow.Keras.Utils;
7+
using static Tensorflow.Binding;
8+
9+
namespace Tensorflow.Keras.Engine
10+
{
11+
public partial class Layer
12+
{
13+
/// <summary>
14+
/// Wraps `call`, applying pre- and post-processing steps.
15+
/// </summary>
16+
/// <param name="input"></param>
17+
/// <param name="state"></param>
18+
/// <param name="is_training"></param>
19+
/// <returns></returns>
20+
public Tensors Apply(Tensors inputs, Tensor state = null, bool is_training = false)
21+
{
22+
callContext = callContext ?? new ThreadLocal<CallContext>()
23+
{
24+
Value = new CallContext()
25+
};
26+
27+
if (_in_functional_construction_mode(inputs))
28+
return FunctionalConstructionCall(inputs);
29+
30+
Tensors outputs = null;
31+
32+
var eager = tf.executing_eagerly();
33+
using var ctxManager = CallContext.enter();
34+
35+
string nameScope = "";
36+
if (eager)
37+
nameScope = Name;
38+
else
39+
nameScope = _name_scope();
40+
41+
if (!inputs.IsEagerTensor)
42+
tf.Context.graph_mode();
43+
44+
tf_with(ops.name_scope(nameScope), scope =>
45+
{
46+
if (!built)
47+
MaybeBuild(inputs);
48+
49+
outputs = call(inputs, state: state, is_training: is_training);
50+
51+
outputs = _set_connectivity_metadata_(inputs, outputs);
52+
_handle_activity_regularization(inputs, outputs);
53+
_set_mask_metadata(inputs, outputs, null);
54+
});
55+
56+
if (!inputs.IsEagerTensor)
57+
tf.Context.restore_mode();
58+
59+
return outputs;
60+
}
61+
}
62+
}
Lines changed: 58 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,58 @@
1+
using System;
2+
using System.Collections.Generic;
3+
using System.Text;
4+
using Tensorflow.Keras.Utils;
5+
using static Tensorflow.Binding;
6+
7+
namespace Tensorflow.Keras.Engine
8+
{
9+
public partial class Layer
10+
{
11+
Tensors FunctionalConstructionCall(Tensors inputs)
12+
{
13+
bool mask_arg_passed_by_framework = false;
14+
bool training_arg_passed_by_framework = false;
15+
Tensor training_value = null;
16+
if (training_value == null)
17+
{
18+
training_arg_passed_by_framework = true;
19+
}
20+
21+
if (base_layer_utils.needs_keras_history(inputs))
22+
base_layer_utils.create_keras_history(inputs);
23+
24+
Tensors outputs = null;
25+
using var ctxManager = CallContext.enter();
26+
27+
// using var graph = tf.keras.backend.get_graph().as_default();
28+
29+
if (!inputs.IsEagerTensor)
30+
tf.Context.graph_mode();
31+
32+
tf_with(ops.name_scope(_name_scope()), scope =>
33+
{
34+
MaybeBuild(inputs);
35+
36+
// Wrapping `call` function in autograph to allow for dynamic control
37+
// flow and control dependencies in call. We are limiting this to
38+
// subclassed layers as autograph is strictly needed only for
39+
// subclassed layers and models.
40+
// tf_convert will respect the value of autograph setting in the
41+
// enclosing tf.function, if any.
42+
if (!dynamic)
43+
throw new NotImplementedException("");
44+
45+
outputs = call(inputs);
46+
47+
outputs = _set_connectivity_metadata_(inputs, outputs);
48+
_handle_activity_regularization(inputs, outputs);
49+
_set_mask_metadata(inputs, outputs, null);
50+
});
51+
52+
if (!inputs.IsEagerTensor)
53+
tf.Context.restore_mode();
54+
55+
return outputs;
56+
}
57+
}
58+
}

0 commit comments

Comments
 (0)

Back | FazBrowse Home | New Git URL