| FazBrowse GitHub Viewer | Trending | | Home |
| Tools: [Download Repo ZIP] [Original HTTPS Page] |
1 parent b79d6bc commit a0ec655
12 files changed
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -11,7 +11,7 @@ public class NodeArgs | |||
| 11 | 11 | public Layer[] InboundLayers { get; set; } | |
| 12 | 12 | public int[] NodeIndices { get; set; } | |
| 13 | 13 | public int[] TensorIndices { get; set; } | |
| 14 | - public Tensor InputTensors { get; set; } | ||
| 14 | + public Tensors InputTensors { get; set; } | ||
| 15 | 15 | public Tensors Outputs { get; set; } | |
| 16 | 16 | } | |
| 17 | 17 | } | |
| Original file line number | Diff line number | Diff 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 | + } | ||
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -4,6 +4,7 @@ | |||
| 4 | 4 | using System.Security.Cryptography.X509Certificates; | |
| 5 | 5 | using System.Text; | |
| 6 | 6 | using Tensorflow.Keras.ArgsDefinition; | |
| 7 | + using Tensorflow.Keras.Utils; | ||
| 7 | 8 | ||
| 8 | 9 | namespace Tensorflow.Keras.Engine | |
| 9 | 10 | { | |
@@ -50,7 +51,7 @@ void _init_graph_network(Tensors inputs, Tensors outputs) | |||
| 50 | 51 | _autocast = false; | |
| 51 | 52 | ||
| 52 | 53 | if (outputs.Any(x => x.KerasHistory == null)) | |
| 53 | - BaseLayerUtils.CreateKerasHistoryHelper(outputs); | ||
| 54 | + base_layer_utils.create_keras_history(outputs); | ||
| 54 | 55 | ||
| 55 | 56 | // Build self._output_layers: | |
| 56 | 57 | foreach (var x in outputs) | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -9,7 +9,7 @@ namespace Tensorflow.Keras.Engine | |||
| 9 | 9 | /// </summary> | |
| 10 | 10 | public class KerasHistory | |
| 11 | 11 | { | |
| 12 | - Layer layer; | ||
| 12 | + public Layer layer; | ||
| 13 | 13 | int node_index; | |
| 14 | 14 | int tensor_index; | |
| 15 | 15 | public Tensor tensor; | |
@@ -20,6 +20,7 @@ public KerasHistory(Layer layer, int node_index, int tensor_index, Tensor tensor | |||
| 20 | 20 | this.node_index = node_index; | |
| 21 | 21 | this.tensor_index = tensor_index; | |
| 22 | 22 | this.tensor = tensor; | |
| 23 | + Layer.KerasHistories.Add(this); | ||
| 23 | 24 | Console.WriteLine(tensor.name); | |
| 24 | 25 | } | |
| 25 | 26 | ||
| Original file line number | Diff line number | Diff 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 | + } | ||
| Original file line number | Diff line number | Diff 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 | + } | ||
| Original file line number | Diff line number | Diff 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 | + } | ||
| Back | FazBrowse Home | New Git URL |
0 commit comments