| FazBrowse GitHub Viewer | Trending | | Home |
| Tools: [Download Repo ZIP] [Original HTTPS Page] |
1 parent 980201d commit 7a706c9
11 files changed
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -69,7 +69,8 @@ public partial class Graph : IPython, IDisposable | |||
| 69 | 69 | private List<Tensor> _unfeedable_tensors = new List<Tensor>(); | |
| 70 | 70 | ||
| 71 | 71 | public string _name_stack = ""; | |
| 72 | - public string _graph_key; | ||
| 72 | + private string _graph_key; | ||
| 73 | + public string graph_key => _graph_key; | ||
| 73 | 74 | public string _last_loss_reduction; | |
| 74 | 75 | ||
| 75 | 76 | public Status Status { get; } | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -19,6 +19,10 @@ public void __enter__() | |||
| 19 | 19 | ||
| 20 | 20 | } | |
| 21 | 21 | ||
| 22 | + /// <summary> | ||
| 23 | + /// Adds a layer instance on top of the layer stack. | ||
| 24 | + /// </summary> | ||
| 25 | + /// <param name="layer"></param> | ||
| 22 | 26 | public void add(Layer layer) | |
| 23 | 27 | { | |
| 24 | 28 | built = false; | |
@@ -32,7 +36,7 @@ public void add(Layer layer) | |||
| 32 | 36 | var x = keras.layers.Input( | |
| 33 | 37 | batch_shape: batch_shape, | |
| 34 | 38 | dtype: dtype, | |
| 35 | - name: layer._name + "_input"); | ||
| 39 | + name: layer.name + "_input"); | ||
| 36 | 40 | ||
| 37 | 41 | // This will build the current layer | |
| 38 | 42 | // and create the node connecting the current layer | |
| 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 static Tensorflow.Python; | ||
| 7 | 8 | ||
| 8 | 9 | namespace Tensorflow.Keras.Layers | |
| 9 | 10 | { | |
@@ -33,7 +34,8 @@ public class Layer : CheckpointableBase | |||
| 33 | 34 | protected InputSpec input_spec; | |
| 34 | 35 | protected bool supports_masking; | |
| 35 | 36 | protected List<RefVariable> _trainable_weights; | |
| 36 | - public string _name; | ||
| 37 | + private string _name; | ||
| 38 | + public string name => _name; | ||
| 37 | 39 | protected string _base_name; | |
| 38 | 40 | protected bool _compute_previous_mask; | |
| 39 | 41 | protected List<Operation> _updates; | |
@@ -85,17 +87,24 @@ public Tensor __call__(Tensor[] inputs, | |||
| 85 | 87 | // Handle Keras mask propagation from previous layer to current layer. | |
| 86 | 88 | Python.with(ops.name_scope(_name_scope()), delegate | |
| 87 | 89 | { | |
| 88 | - if (!built) | ||
| 90 | + /*if (!built) | ||
| 89 | 91 | { | |
| 90 | 92 | _maybe_build(inputs); | |
| 91 | 93 | built = true; | |
| 92 | - } | ||
| 94 | + }*/ | ||
| 93 | 95 | ||
| 94 | 96 | if (build_graph) | |
| 95 | 97 | { | |
| 96 | 98 | // Symbolic execution on symbolic tensors. We will attempt to build | |
| 97 | 99 | // the corresponding TF subgraph inside `backend.get_graph()` | |
| 98 | - var graph = 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 | + | ||
| 99 | 108 | outputs = call(inputs[0], training: training); | |
| 100 | 109 | _handle_activity_regularization(inputs[0], outputs); | |
| 101 | 110 | _set_mask_metadata(inputs[0], outputs, null); | |
@@ -130,13 +139,17 @@ protected virtual Tensor call(Tensor inputs, Tensor training = null) | |||
| 130 | 139 | ||
| 131 | 140 | protected virtual string _name_scope() | |
| 132 | 141 | { | |
| 133 | - return null; | ||
| 142 | + return name; | ||
| 134 | 143 | } | |
| 135 | 144 | ||
| 136 | - protected void _maybe_build(Tensor[] inputs) | ||
| 145 | + protected void _maybe_build(Tensor input) | ||
| 137 | 146 | { | |
| 138 | - var input_list = inputs; | ||
| 139 | - build(input_list[0].GetShape()); | ||
| 147 | + // Check input assumptions set before layer building, e.g. input rank. | ||
| 148 | + if (built) | ||
| 149 | + return; | ||
| 150 | + | ||
| 151 | + build(input.GetShape()); | ||
| 152 | + built = true; | ||
| 140 | 153 | } | |
| 141 | 154 | ||
| 142 | 155 | protected virtual void build(TensorShape input_shape) | |
@@ -160,7 +173,7 @@ protected virtual RefVariable add_weight(string name, | |||
| 160 | 173 | var variable = _add_variable_with_custom_getter(name, | |
| 161 | 174 | shape, | |
| 162 | 175 | dtype: dtype, | |
| 163 | - getter: getter == null ? base_layer_utils.make_variable : getter, | ||
| 176 | + //getter: getter == null ? base_layer_utils.make_variable : getter, | ||
| 164 | 177 | overwrite: true, | |
| 165 | 178 | initializer: initializer, | |
| 166 | 179 | trainable: trainable.Value); | |
@@ -176,12 +189,12 @@ protected virtual void add_update(Tensor[] updates, bool inputs = false) | |||
| 176 | 189 | _updates.AddRange(updates_op); | |
| 177 | 190 | } | |
| 178 | 191 | ||
| 179 | - protected virtual void _init_set_name(string name) | ||
| 192 | + protected virtual void _init_set_name(string name, bool zero_based = true) | ||
| 180 | 193 | { | |
| 181 | - string base_name = name; | ||
| 182 | 194 | if (name == null) | |
| 183 | - (_name, base_name) = _make_unique_name(); | ||
| 184 | - _base_name = base_name; | ||
| 195 | + _name = base_layer_utils.unique_layer_name(generic_utils.to_snake_case(this.GetType().Name), zero_based: zero_based); | ||
| 196 | + else | ||
| 197 | + _name = name; | ||
| 185 | 198 | } | |
| 186 | 199 | ||
| 187 | 200 | protected virtual (string, string) _make_unique_name() | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -30,26 +30,6 @@ public NDArray pad_sequences(NDArray sequences, | |||
| 30 | 30 | object value = null) | |
| 31 | 31 | { | |
| 32 | 32 | int[] length = new int[sequences.size]; | |
| 33 | - switch (sequences.dtype.Name) | ||
| 34 | - { | ||
| 35 | - case "Object": | ||
| 36 | - for (int i = 0; i < sequences.size; i++) | ||
| 37 | - { | ||
| 38 | - switch (sequences.Data<object>(i)) | ||
| 39 | - { | ||
| 40 | - case string data: | ||
| 41 | - length[i] = Regex.Matches(data, ",").Count; | ||
| 42 | - break; | ||
| 43 | - } | ||
| 44 | - } | ||
| 45 | - break; | ||
| 46 | - case "Int32": | ||
| 47 | - for (int i = 0; i < sequences.size; i++) | ||
| 48 | - length[i] = Regex.Matches(sequences.Data<object>(i).ToString(), ",").Count; | ||
| 49 | - break; | ||
| 50 | - default: | ||
| 51 | - throw new NotImplementedException($"pad_sequences: {sequences.dtype.Name}"); | ||
| 52 | - } | ||
| 53 | 33 | ||
| 54 | 34 | if (maxlen == null) | |
| 55 | 35 | maxlen = length.Max(); | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -1,35 +1,96 @@ | |||
| 1 | 1 | using System; | |
| 2 | 2 | using System.Collections.Generic; | |
| 3 | + using System.Linq; | ||
| 3 | 4 | using System.Text; | |
| 5 | + using static Tensorflow.Python; | ||
| 4 | 6 | ||
| 5 | 7 | namespace Tensorflow.Keras.Utils | |
| 6 | 8 | { | |
| 7 | 9 | public class base_layer_utils | |
| 8 | 10 | { | |
| 11 | + /// <summary> | ||
| 12 | + /// Adds a new variable to the layer. | ||
| 13 | + /// </summary> | ||
| 14 | + /// <param name="name"></param> | ||
| 15 | + /// <param name="shape"></param> | ||
| 16 | + /// <param name="dtype"></param> | ||
| 17 | + /// <param name="initializer"></param> | ||
| 18 | + /// <param name="trainable"></param> | ||
| 19 | + /// <returns></returns> | ||
| 9 | 20 | public static RefVariable make_variable(string name, | |
| 10 | 21 | int[] shape, | |
| 11 | 22 | TF_DataType dtype = TF_DataType.TF_FLOAT, | |
| 12 | 23 | IInitializer initializer = null, | |
| 13 | - bool trainable = false) | ||
| 24 | + bool trainable = true, | ||
| 25 | + bool use_resource = true) | ||
| 14 | 26 | { | |
| 15 | - throw new NotImplementedException(""); | ||
| 27 | + var initializing_from_value = false; | ||
| 28 | + | ||
| 29 | + ops.init_scope(); | ||
| 30 | + | ||
| 31 | + Func<Tensor> init_val = ()=> initializer.call(new TensorShape(shape), dtype: dtype); | ||
| 32 | + | ||
| 33 | + var variable_dtype = dtype.as_base_dtype(); | ||
| 34 | + var v = tf.Variable(init_val); | ||
| 35 | + | ||
| 36 | + return v; | ||
| 16 | 37 | } | |
| 17 | 38 | ||
| 18 | 39 | /// <summary> | |
| 19 | 40 | /// Makes a layer name (or arbitrary string) unique within a TensorFlow graph. | |
| 20 | 41 | /// </summary> | |
| 21 | 42 | /// <param name="name"></param> | |
| 22 | 43 | /// <returns></returns> | |
| 23 | - public static string unique_layer_name(string name) | ||
| 44 | + public static string unique_layer_name(string name, Dictionary<(string, string), int> name_uid_map = null, | ||
| 45 | + string[] avoid_names = null, string @namespace = "", bool zero_based = false) | ||
| 24 | 46 | { | |
| 25 | - int number = get_default_graph_uid_map(); | ||
| 26 | - return $"{name}_{number}"; | ||
| 47 | + if(name_uid_map == null) | ||
| 48 | + name_uid_map = get_default_graph_uid_map(); | ||
| 49 | + if (avoid_names == null) | ||
| 50 | + avoid_names = new string[0]; | ||
| 51 | + | ||
| 52 | + string proposed_name = null; | ||
| 53 | + while(proposed_name == null || avoid_names.Contains(proposed_name)) | ||
| 54 | + { | ||
| 55 | + var name_key = (@namespace, name); | ||
| 56 | + if (!name_uid_map.ContainsKey(name_key)) | ||
| 57 | + name_uid_map[name_key] = 0; | ||
| 58 | + | ||
| 59 | + if (zero_based) | ||
| 60 | + { | ||
| 61 | + int number = name_uid_map[name_key]; | ||
| 62 | + if (number > 0) | ||
| 63 | + proposed_name = $"{name}_{number}"; | ||
| 64 | + else | ||
| 65 | + proposed_name = name; | ||
| 66 | + | ||
| 67 | + name_uid_map[name_key] += 1; | ||
| 68 | + } | ||
| 69 | + else | ||
| 70 | + { | ||
| 71 | + name_uid_map[name_key] += 1; | ||
| 72 | + proposed_name = $"{name}_{name_uid_map[name_key]}"; | ||
| 73 | + } | ||
| 74 | + } | ||
| 75 | + | ||
| 76 | + return proposed_name; | ||
| 27 | 77 | } | |
| 28 | 78 | ||
| 29 | - public static int get_default_graph_uid_map() | ||
| 79 | + public static Dictionary<(string, string), int> get_default_graph_uid_map() | ||
| 30 | 80 | { | |
| 31 | 81 | var graph = ops.get_default_graph(); | |
| 32 | - return graph._next_id(); | ||
| 82 | + Dictionary<(string, string), int> name_uid_map = null; | ||
| 83 | + if (backend.PER_GRAPH_LAYER_NAME_UIDS.ContainsKey(graph.graph_key)) | ||
| 84 | + { | ||
| 85 | + name_uid_map = backend.PER_GRAPH_LAYER_NAME_UIDS[graph.graph_key]; | ||
| 86 | + } | ||
| 87 | + else | ||
| 88 | + { | ||
| 89 | + name_uid_map = new Dictionary<(string, string), int>(); | ||
| 90 | + backend.PER_GRAPH_LAYER_NAME_UIDS[graph.graph_key] = name_uid_map; | ||
| 91 | + } | ||
| 92 | + | ||
| 93 | + return name_uid_map; | ||
| 33 | 94 | } | |
| 34 | 95 | } | |
| 35 | 96 | } | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -6,6 +6,13 @@ namespace Tensorflow.Keras | |||
| 6 | 6 | { | |
| 7 | 7 | public class backend | |
| 8 | 8 | { | |
| 9 | + /// <summary> | ||
| 10 | + /// A global dictionary mapping graph objects to an index of counters used | ||
| 11 | + /// for various layer names in each graph. | ||
| 12 | + /// Allows to give unique autogenerated names to layers, in a graph-specific way. | ||
| 13 | + /// </summary> | ||
| 14 | + public static Dictionary<string, Dictionary<(string, string), int>> PER_GRAPH_LAYER_NAME_UIDS = new Dictionary<string, Dictionary<(string, string), int>>(); | ||
| 15 | + | ||
| 9 | 16 | public static void track_variable(RefVariable v) | |
| 10 | 17 | { | |
| 11 | 18 | ||
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -1,34 +1,12 @@ | |||
| 1 | 1 | using System; | |
| 2 | 2 | using System.Collections.Generic; | |
| 3 | 3 | using System.Text; | |
| 4 | + using Tensorflow.Train; | ||
| 4 | 5 | ||
| 5 | 6 | namespace Tensorflow | |
| 6 | 7 | { | |
| 7 | - public abstract class CheckpointableBase | ||
| 8 | + public abstract class CheckpointableBase : Trackable | ||
| 8 | 9 | { | |
| 9 | - /// <summary> | ||
| 10 | - /// Restore-on-create for a variable be saved with this `Checkpointable`. | ||
| 11 | - /// </summary> | ||
| 12 | - /// <returns></returns> | ||
| 13 | - protected virtual RefVariable _add_variable_with_custom_getter(string name, | ||
| 14 | - int[] shape, | ||
| 15 | - TF_DataType dtype = TF_DataType.TF_FLOAT, | ||
| 16 | - IInitializer initializer = null, | ||
| 17 | - Func<string, int[], TF_DataType, IInitializer, bool, RefVariable> getter = null, | ||
| 18 | - bool overwrite = false, | ||
| 19 | - bool trainable = false) | ||
| 20 | - { | ||
| 21 | - var new_variable = getter(name, shape, dtype, initializer, trainable); | ||
| 22 | - if (!overwrite || new_variable is RefVariable) | ||
| 23 | - return _track_checkpointable(new_variable, name: name, | ||
| 24 | - overwrite: overwrite); | ||
| 25 | - else | ||
| 26 | - return new_variable; | ||
| 27 | - } | ||
| 28 | 10 | ||
| 29 | - protected RefVariable _track_checkpointable(RefVariable checkpointable, string name, bool overwrite = false) | ||
| 30 | - { | ||
| 31 | - return checkpointable; | ||
| 32 | - } | ||
| 33 | 11 | } | |
| 34 | 12 | } | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -0,0 +1,34 @@ | |||
| 1 | + using System; | ||
| 2 | + using System.Collections.Generic; | ||
| 3 | + using System.Text; | ||
| 4 | + | ||
| 5 | + namespace Tensorflow.Train | ||
| 6 | + { | ||
| 7 | + public abstract class Trackable | ||
| 8 | + { | ||
| 9 | + /// <summary> | ||
| 10 | + /// Restore-on-create for a variable be saved with this `Checkpointable`. | ||
| 11 | + /// </summary> | ||
| 12 | + /// <returns></returns> | ||
| 13 | + protected virtual RefVariable _add_variable_with_custom_getter(string name, | ||
| 14 | + int[] shape, | ||
| 15 | + TF_DataType dtype = TF_DataType.TF_FLOAT, | ||
| 16 | + IInitializer initializer = null, | ||
| 17 | + Func<string, int[], TF_DataType, IInitializer, bool, RefVariable> getter = null, | ||
| 18 | + bool overwrite = false, | ||
| 19 | + bool trainable = false) | ||
| 20 | + { | ||
| 21 | + var new_variable = getter(name, shape, dtype, initializer, trainable); | ||
| 22 | + if (!overwrite || new_variable is RefVariable) | ||
| 23 | + return _track_checkpointable(new_variable, name: name, | ||
| 24 | + overwrite: overwrite); | ||
| 25 | + else | ||
| 26 | + return new_variable; | ||
| 27 | + } | ||
| 28 | + | ||
| 29 | + protected RefVariable _track_checkpointable(RefVariable checkpointable, string name, bool overwrite = false) | ||
| 30 | + { | ||
| 31 | + return checkpointable; | ||
| 32 | + } | ||
| 33 | + } | ||
| 34 | + } | ||
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -111,7 +111,7 @@ private void _init_from_args(object initial_value, | |||
| 111 | 111 | ||
| 112 | 112 | // Store the graph key so optimizers know how to only retrieve variables from | |
| 113 | 113 | // this graph. | |
| 114 | - _graph_key = ops.get_default_graph()._graph_key; | ||
| 114 | + _graph_key = ops.get_default_graph().graph_key; | ||
| 115 | 115 | ||
| 116 | 116 | _trainable = trainable; | |
| 117 | 117 | if (trainable && !collections.Contains(ops.GraphKeys.TRAINABLE_VARIABLES)) | |
| Back | FazBrowse Home | New Git URL |
0 commit comments