| FazBrowse GitHub Viewer | Trending | | Home |
| Tools: [Download Repo ZIP] [Original HTTPS Page] |
1 parent a5ae56a commit d7c7d3d
11 files changed
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -14,15 +14,15 @@ public static partial class tf | |||
| 14 | 14 | ||
| 15 | 15 | public static variable_scope variable_scope(string name, | |
| 16 | 16 | string default_name = null, | |
| 17 | - object values = null, | ||
| 17 | + Tensor[] values = null, | ||
| 18 | 18 | bool auxiliary_name_scope = true) => new variable_scope(name, | |
| 19 | 19 | default_name, | |
| 20 | 20 | values, | |
| 21 | 21 | auxiliary_name_scope); | |
| 22 | 22 | ||
| 23 | 23 | public static variable_scope variable_scope(VariableScope scope, | |
| 24 | 24 | string default_name = null, | |
| 25 | - object values = null, | ||
| 25 | + Tensor[] values = null, | ||
| 26 | 26 | bool? reuse = null, | |
| 27 | 27 | bool auxiliary_name_scope = true) => new variable_scope(scope, | |
| 28 | 28 | default_name, | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -4,6 +4,7 @@ | |||
| 4 | 4 | using System.Text; | |
| 5 | 5 | using Tensorflow.Keras.Engine; | |
| 6 | 6 | using Tensorflow.Keras.Utils; | |
| 7 | + using Tensorflow.Train; | ||
| 7 | 8 | using static Tensorflow.Python; | |
| 8 | 9 | ||
| 9 | 10 | namespace Tensorflow.Keras.Layers | |
@@ -14,7 +15,7 @@ namespace Tensorflow.Keras.Layers | |||
| 14 | 15 | /// as convolution, batch norm, etc. These operations require managing weights, | |
| 15 | 16 | /// losses, updates, and inter-layer connectivity. | |
| 16 | 17 | /// </summary> | |
| 17 | - public class Layer : CheckpointableBase | ||
| 18 | + public class Layer : AutoTrackable | ||
| 18 | 19 | { | |
| 19 | 20 | /// <summary> | |
| 20 | 21 | /// Indicates whether `build` needs to be called upon layer call, to create | |
@@ -84,32 +85,35 @@ public Tensor __call__(Tensor[] inputs, | |||
| 84 | 85 | // models using the functional API). | |
| 85 | 86 | bool build_graph = tf_utils.are_all_symbolic_tensors(input_list); | |
| 86 | 87 | ||
| 87 | - // Handle Keras mask propagation from previous layer to current layer. | ||
| 88 | - Python.with(ops.name_scope(_name_scope()), delegate | ||
| 88 | + if (build_graph) | ||
| 89 | 89 | { | |
| 90 | - /*if (!built) | ||
| 91 | - { | ||
| 92 | - _maybe_build(inputs); | ||
| 93 | - built = true; | ||
| 94 | - }*/ | ||
| 90 | + // Only create Keras history if at least one tensor originates from a | ||
| 91 | + // `keras.Input`. Otherwise this Layer may be being used outside the Keras | ||
| 92 | + // framework. | ||
| 93 | + // base_layer_utils.create_keras_history(inputs) | ||
| 94 | + } | ||
| 95 | 95 | ||
| 96 | - if (build_graph) | ||
| 96 | + // with base_layer_utils.call_context(self): | ||
| 97 | + | ||
| 98 | + // Handle Keras mask propagation from previous layer to current layer. | ||
| 99 | + // with base_layer_utils.call_context(self): | ||
| 100 | + // Check input assumptions set after layer building, e.g. input shape. | ||
| 101 | + if (build_graph) | ||
| 102 | + { | ||
| 103 | + // Symbolic execution on symbolic tensors. We will attempt to build | ||
| 104 | + // the corresponding TF subgraph inside `backend.get_graph()` | ||
| 105 | + var graph = backend.get_graph().as_default(); | ||
| 106 | + with(ops.name_scope(_name_scope()), delegate | ||
| 97 | 107 | { | |
| 98 | - // Symbolic execution on symbolic tensors. We will attempt to build | ||
| 99 | - // the corresponding TF subgraph inside `backend.get_graph()` | ||
| 100 | - var graph = backend.get_graph().as_default(); | ||
| 101 | - with(ops.name_scope(_name_scope()), delegate | ||
| 102 | - { | ||
| 103 | - // Build layer if applicable (if the `build` method has been | ||
| 104 | - // overridden). | ||
| 105 | - _maybe_build(inputs[0]); | ||
| 106 | - }); | ||
| 107 | - | ||
| 108 | - outputs = call(inputs[0], training: training); | ||
| 109 | - _handle_activity_regularization(inputs[0], outputs); | ||
| 110 | - _set_mask_metadata(inputs[0], outputs, null); | ||
| 111 | - } | ||
| 112 | - }); | ||
| 108 | + // Build layer if applicable (if the `build` method has been | ||
| 109 | + // overridden). | ||
| 110 | + _maybe_build(inputs[0]); | ||
| 111 | + }); | ||
| 112 | + | ||
| 113 | + outputs = call(inputs[0], training: training); | ||
| 114 | + _handle_activity_regularization(inputs[0], outputs); | ||
| 115 | + _set_mask_metadata(inputs[0], outputs, null); | ||
| 116 | + } | ||
| 113 | 117 | ||
| 114 | 118 | return outputs; | |
| 115 | 119 | } | |
@@ -147,6 +151,8 @@ protected void _maybe_build(Tensor input) | |||
| 147 | 151 | // Check input assumptions set before layer building, e.g. input rank. | |
| 148 | 152 | if (built) | |
| 149 | 153 | return; | |
| 154 | + if (_dtype == TF_DataType.DtInvalid) | ||
| 155 | + _dtype = input.dtype; | ||
| 150 | 156 | ||
| 151 | 157 | build(input.GetShape()); | |
| 152 | 158 | built = true; | |
@@ -170,10 +176,21 @@ protected virtual RefVariable add_weight(string name, | |||
| 170 | 176 | if (trainable == null) | |
| 171 | 177 | trainable = true; | |
| 172 | 178 | ||
| 179 | + // Initialize variable when no initializer provided | ||
| 180 | + if(initializer == null) | ||
| 181 | + { | ||
| 182 | + // If dtype is DT_FLOAT, provide a uniform unit scaling initializer | ||
| 183 | + if (dtype.is_floating()) | ||
| 184 | + initializer = tf.glorot_uniform_initializer; | ||
| 185 | + else if (dtype.is_integer()) | ||
| 186 | + initializer = tf.zeros_initializer; | ||
| 187 | + else | ||
| 188 | + throw new ValueError($"An initializer for variable {name} of type {dtype.as_base_dtype()} is required for layer {this.name}"); | ||
| 189 | + } | ||
| 173 | 190 | var variable = _add_variable_with_custom_getter(name, | |
| 174 | 191 | shape, | |
| 175 | 192 | dtype: dtype, | |
| 176 | - //getter: getter == null ? base_layer_utils.make_variable : getter, | ||
| 193 | + getter: getter, // getter == null ? base_layer_utils.make_variable : getter, | ||
| 177 | 194 | overwrite: true, | |
| 178 | 195 | initializer: initializer, | |
| 179 | 196 | trainable: trainable.Value); | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -12,10 +12,11 @@ public class backend | |||
| 12 | 12 | /// Allows to give unique autogenerated names to layers, in a graph-specific way. | |
| 13 | 13 | /// </summary> | |
| 14 | 14 | public static Dictionary<string, Dictionary<(string, string), int>> PER_GRAPH_LAYER_NAME_UIDS = new Dictionary<string, Dictionary<(string, string), int>>(); | |
| 15 | - | ||
| 15 | + public static Dictionary<string, RefVariable> _GRAPH_VARIABLES = new Dictionary<string, RefVariable>(); | ||
| 16 | 16 | public static void track_variable(RefVariable v) | |
| 17 | 17 | { | |
| 18 | - | ||
| 18 | + var graph = v.graph; | ||
| 19 | + _GRAPH_VARIABLES[graph.graph_key] = v; | ||
| 19 | 20 | } | |
| 20 | 21 | ||
| 21 | 22 | public static Tensor placeholder(int[] shape = null, | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -51,9 +51,14 @@ public Tensor __call__(Tensor inputs, | |||
| 51 | 51 | auxiliary_name_scope: false); | |
| 52 | 52 | } | |
| 53 | 53 | ||
| 54 | - with(scope_context_manager, scope2 => _current_scope = scope2); | ||
| 55 | - // Actually call layer | ||
| 56 | - var outputs = base.__call__(new Tensor[] { inputs }, training: training); | ||
| 54 | + Tensor outputs = null; | ||
| 55 | + with(scope_context_manager, scope2 => | ||
| 56 | + { | ||
| 57 | + _current_scope = scope2; | ||
| 58 | + // Actually call layer | ||
| 59 | + outputs = base.__call__(new Tensor[] { inputs }, training: training); | ||
| 60 | + }); | ||
| 61 | + | ||
| 57 | 62 | ||
| 58 | 63 | // Update global default collections. | |
| 59 | 64 | _add_elements_to_collection(_updates.ToArray(), new string[] { ops.GraphKeys.UPDATE_OPS }); | |
@@ -80,6 +85,17 @@ protected virtual void _add_elements_to_collection(Operation[] elements, string[ | |||
| 80 | 85 | } | |
| 81 | 86 | } | |
| 82 | 87 | ||
| 88 | + /// <summary> | ||
| 89 | + /// Adds a new variable to the layer, or gets an existing one; returns it. | ||
| 90 | + /// </summary> | ||
| 91 | + /// <param name="name"></param> | ||
| 92 | + /// <param name="shape"></param> | ||
| 93 | + /// <param name="dtype"></param> | ||
| 94 | + /// <param name="initializer"></param> | ||
| 95 | + /// <param name="trainable"></param> | ||
| 96 | + /// <param name="synchronization"></param> | ||
| 97 | + /// <param name="aggregation"></param> | ||
| 98 | + /// <returns></returns> | ||
| 83 | 99 | protected virtual RefVariable add_weight(string name, | |
| 84 | 100 | int[] shape, | |
| 85 | 101 | TF_DataType dtype = TF_DataType.DtInvalid, | |
@@ -157,7 +173,10 @@ private void _set_scope(VariableScope scope = null) | |||
| 157 | 173 | else | |
| 158 | 174 | { | |
| 159 | 175 | with(tf.variable_scope(scope, default_name: _base_name), | |
| 160 | - captured_scope => _scope = captured_scope); | ||
| 176 | + captured_scope => | ||
| 177 | + { | ||
| 178 | + _scope = captured_scope; | ||
| 179 | + }); | ||
| 161 | 180 | } | |
| 162 | 181 | ||
| 163 | 182 | } | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -0,0 +1,10 @@ | |||
| 1 | + using System; | ||
| 2 | + using System.Collections.Generic; | ||
| 3 | + using System.Text; | ||
| 4 | + | ||
| 5 | + namespace Tensorflow.Train | ||
| 6 | + { | ||
| 7 | + public abstract class AutoTrackable : Trackable | ||
| 8 | + { | ||
| 9 | + } | ||
| 10 | + } | ||
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -5,7 +5,7 @@ | |||
| 5 | 5 | ||
| 6 | 6 | namespace Tensorflow | |
| 7 | 7 | { | |
| 8 | - public abstract class CheckpointableBase : Trackable | ||
| 8 | + public abstract class CheckpointableBase : AutoTrackable | ||
| 9 | 9 | { | |
| 10 | 10 | ||
| 11 | 11 | } | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -18,7 +18,13 @@ protected virtual RefVariable _add_variable_with_custom_getter(string name, | |||
| 18 | 18 | bool overwrite = false, | |
| 19 | 19 | bool trainable = false) | |
| 20 | 20 | { | |
| 21 | + var checkpoint_initializer = true; | ||
| 21 | 22 | var new_variable = getter(name, shape, dtype, initializer, trainable); | |
| 23 | + | ||
| 24 | + // If we set an initializer and the variable processed it, tracking will not | ||
| 25 | + // assign again. It will add this variable to our dependencies, and if there | ||
| 26 | + // is a non-trivial restoration queued, it will handle that. This also | ||
| 27 | + // handles slot variables. | ||
| 22 | 28 | if (!overwrite || new_variable is RefVariable) | |
| 23 | 29 | return _track_checkpointable(new_variable, name: name, | |
| 24 | 30 | overwrite: overwrite); | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -35,7 +35,7 @@ public PureVariableScope(VariableScope scope, | |||
| 35 | 35 | _old_name_scope = old_name_scope; | |
| 36 | 36 | _var_store = variable_scope._get_default_variable_store(); | |
| 37 | 37 | _var_scope_store = variable_scope.get_variable_scope_store(); | |
| 38 | - _new_name = _scope._name; | ||
| 38 | + _new_name = _scope.name; | ||
| 39 | 39 | ||
| 40 | 40 | string name_scope = _scope._name_scope; | |
| 41 | 41 | variable_scope_object = new VariableScope(_reuse, | |
@@ -55,7 +55,7 @@ public void __enter__() | |||
| 55 | 55 | } | |
| 56 | 56 | else | |
| 57 | 57 | { | |
| 58 | - _new_name = string.IsNullOrEmpty(_old._name) ? _name : _old._name + "/" + _name; | ||
| 58 | + _new_name = string.IsNullOrEmpty(_old.name) ? _name : _old.name + "/" + _name; | ||
| 59 | 59 | _reuse = _reuse || _old.resue; | |
| 60 | 60 | string name_scope = _old_name_scope == null ? _name : _old_name_scope; | |
| 61 | 61 | ||
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -15,7 +15,8 @@ public class VariableScope | |||
| 15 | 15 | public bool resue; | |
| 16 | 16 | ||
| 17 | 17 | private TF_DataType _dtype; | |
| 18 | - public string _name { get; set; } | ||
| 18 | + string _name; | ||
| 19 | + public string name => _name; | ||
| 19 | 20 | public string _name_scope { get; set; } | |
| 20 | 21 | public string original_name_scope => _name_scope; | |
| 21 | 22 | ||
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -19,15 +19,17 @@ public class variable_scope : IPython | |||
| 19 | 19 | private string _name; | |
| 20 | 20 | private VariableScope _scope; | |
| 21 | 21 | private string _default_name; | |
| 22 | - private object _values; | ||
| 22 | + private Tensor[] _values; | ||
| 23 | 23 | private ops.NameScope _current_name_scope; | |
| 24 | 24 | private bool _auxiliary_name_scope; | |
| 25 | 25 | private PureVariableScope _cached_pure_variable_scope; | |
| 26 | 26 | private bool? _reuse; | |
| 27 | + bool _in_graph_mode; | ||
| 28 | + protected Graph _graph; | ||
| 27 | 29 | ||
| 28 | 30 | public variable_scope(string name, | |
| 29 | - string default_name = "", | ||
| 30 | - object values = null, | ||
| 31 | + string default_name = "", | ||
| 32 | + Tensor[] values = null, | ||
| 31 | 33 | bool? reuse = null, | |
| 32 | 34 | bool auxiliary_name_scope = true) | |
| 33 | 35 | { | |
@@ -45,7 +47,7 @@ public variable_scope(string name, | |||
| 45 | 47 | ||
| 46 | 48 | public variable_scope(VariableScope scope, | |
| 47 | 49 | string default_name = "", | |
| 48 | - object values = null, | ||
| 50 | + Tensor[] values = null, | ||
| 49 | 51 | bool? reuse = null, | |
| 50 | 52 | bool auxiliary_name_scope = true) | |
| 51 | 53 | { | |
@@ -58,6 +60,11 @@ public variable_scope(VariableScope scope, | |||
| 58 | 60 | if (_default_name == null && _scope == null) | |
| 59 | 61 | throw new TypeError("If default_name is None then scope is required"); | |
| 60 | 62 | ||
| 63 | + if (_values == null) | ||
| 64 | + _values = new Tensor[0]; | ||
| 65 | + _in_graph_mode = true; | ||
| 66 | + if (_in_graph_mode) | ||
| 67 | + _graph = ops._get_graph_from_inputs(_values); | ||
| 61 | 68 | _auxiliary_name_scope = auxiliary_name_scope; | |
| 62 | 69 | } | |
| 63 | 70 | ||
@@ -87,7 +94,7 @@ private VariableScope _enter_scope_uncached() | |||
| 87 | 94 | ||
| 88 | 95 | if (_name != null || _scope != null) | |
| 89 | 96 | { | |
| 90 | - var name_scope = _name == null ? _scope._name.Split('/').Last() : _name; | ||
| 97 | + var name_scope = _name == null ? _scope.name.Split('/').Last() : _name; | ||
| 91 | 98 | if (name_scope != null || current_name_scope != null) | |
| 92 | 99 | current_name_scope = ops.name_scope(name_scope); | |
| 93 | 100 | current_name_scope.__enter__(); | |
@@ -124,7 +131,7 @@ public static string _get_unique_variable_scope(string prefix) | |||
| 124 | 131 | { | |
| 125 | 132 | var var_scope_store = get_variable_scope_store(); | |
| 126 | 133 | var current_scope = get_variable_scope(); | |
| 127 | - string name = !string.IsNullOrEmpty(current_scope._name) ? current_scope._name + "/" + prefix : prefix; | ||
| 134 | + string name = !string.IsNullOrEmpty(current_scope.name) ? current_scope.name + "/" + prefix : prefix; | ||
| 128 | 135 | if (var_scope_store.variable_scope_count(name) == 0) | |
| 129 | 136 | return prefix; | |
| 130 | 137 | throw new NotImplementedException("_get_unique_variable_scope"); | |
| Back | FazBrowse Home | New Git URL |
0 commit comments