| FazBrowse GitHub Viewer | Trending | | Home |
| Tools: [Download Repo ZIP] [Original HTTPS Page] |
1 parent 71e1fe6 commit bcb803d
10 files changed
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -91,9 +91,18 @@ private T _as_graph_element_locked<T>(T obj, bool allow_tensor = true, bool allo | |||
| 91 | 91 | throw new Exception($"Can not convert a {typeof(T).Name} into a {types_str}."); | |
| 92 | 92 | } | |
| 93 | 93 | ||
| 94 | - public void add_to_collection(string name, object value) | ||
| 94 | + public void add_to_collection<T>(string name, T value) | ||
| 95 | 95 | { | |
| 96 | - _collections[name] = value; | ||
| 96 | + if (_collections.ContainsKey(name)) | ||
| 97 | + (_collections[name] as List<T>).Add(value); | ||
| 98 | + else | ||
| 99 | + _collections[name] = new List<T> { value }; | ||
| 100 | + } | ||
| 101 | + | ||
| 102 | + public void add_to_collections<T>(List<string> names, T value) | ||
| 103 | + { | ||
| 104 | + foreach (string name in names) | ||
| 105 | + add_to_collection(name, value); | ||
| 97 | 106 | } | |
| 98 | 107 | ||
| 99 | 108 | public unsafe Operation create_op(string op_type, List<Tensor> inputs, TF_DataType[] dtypes, | |
@@ -236,9 +245,9 @@ public Operation[] get_operations() | |||
| 236 | 245 | return _nodes_by_name.Values.Select(x => x).ToArray(); | |
| 237 | 246 | } | |
| 238 | 247 | ||
| 239 | - public Dictionary<string, object> get_collection(string name) | ||
| 248 | + public object get_collection(string name) | ||
| 240 | 249 | { | |
| 241 | - return _collections; | ||
| 250 | + return _collections.ContainsKey(name) ? _collections[name] : null; | ||
| 242 | 251 | } | |
| 243 | 252 | ||
| 244 | 253 | public void Dispose() | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -20,7 +20,7 @@ public Operation _apply_op_helper(string op_type_name, string name = "", Diction | |||
| 20 | 20 | name = op_type_name; | |
| 21 | 21 | } | |
| 22 | 22 | ||
| 23 | - string scope = g.unique_name(name) + "/"; | ||
| 23 | + string scope = new ops.name_scope(name); | ||
| 24 | 24 | ||
| 25 | 25 | var default_type_attr_map = new Dictionary<string, object>(); | |
| 26 | 26 | foreach (var attr_def in op_def.Attr) | |
@@ -88,15 +88,22 @@ public Operation _apply_op_helper(string op_type_name, string name = "", Diction | |||
| 88 | 88 | ||
| 89 | 89 | switch (attr_def.Type) | |
| 90 | 90 | { | |
| 91 | + case "string": | ||
| 92 | + attr_value.S = Google.Protobuf.ByteString.CopyFromUtf8((string)value); | ||
| 93 | + break; | ||
| 91 | 94 | case "type": | |
| 92 | 95 | attr_value.Type = _MakeType((TF_DataType)value, attr_def); | |
| 93 | 96 | break; | |
| 94 | 97 | case "bool": | |
| 95 | 98 | attr_value.B = (bool)value; | |
| 96 | 99 | break; | |
| 97 | 100 | case "shape": | |
| 98 | - attr_value.Shape = new TensorShapeProto(); | ||
| 101 | + attr_value.Shape = value == null ? | ||
| 102 | + attr_def.DefaultValue.Shape : | ||
| 103 | + tensor_util.as_shape((long[])value); | ||
| 99 | 104 | break; | |
| 105 | + default: | ||
| 106 | + throw new InvalidDataException($"attr_def.Type {attr_def.Type}"); | ||
| 100 | 107 | } | |
| 101 | 108 | ||
| 102 | 109 | attr_protos[key] = attr_value; | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -73,6 +73,22 @@ public static NDArray convert_to_numpy_ndarray(object values) | |||
| 73 | 73 | return nd; | |
| 74 | 74 | } | |
| 75 | 75 | ||
| 76 | + public static TensorShapeProto as_shape(long[] dims) | ||
| 77 | + { | ||
| 78 | + TensorShapeProto shape = new TensorShapeProto(); | ||
| 79 | + | ||
| 80 | + for (int i = 0; i < dims.Length; i++) | ||
| 81 | + { | ||
| 82 | + var dim = new TensorShapeProto.Types.Dim(); | ||
| 83 | + dim.Size = dims[i]; | ||
| 84 | + dim.Name = $"dim_{i}"; | ||
| 85 | + | ||
| 86 | + shape.Dim.Add(dim); | ||
| 87 | + } | ||
| 88 | + | ||
| 89 | + return shape; | ||
| 90 | + } | ||
| 91 | + | ||
| 76 | 92 | public static TensorShape as_shape(this IShape shape, int[] dims) | |
| 77 | 93 | { | |
| 78 | 94 | return new TensorShape(dims); | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -30,9 +30,14 @@ public Optimizer(double learning_rate, bool use_locking, string name = "") | |||
| 30 | 30 | /// </summary> | |
| 31 | 31 | /// <param name="loss"></param> | |
| 32 | 32 | /// <returns></returns> | |
| 33 | - public Optimizer minimize(Tensor loss, GateGradientType gate_gradients = GateGradientType.GATE_OP) | ||
| 33 | + public Optimizer minimize(Tensor loss, | ||
| 34 | + GateGradientType gate_gradients = GateGradientType.GATE_OP, | ||
| 35 | + bool colocate_gradients_with_ops = false) | ||
| 34 | 36 | { | |
| 35 | - compute_gradients(loss, gate_gradients); | ||
| 37 | + compute_gradients(loss, | ||
| 38 | + gate_gradients: gate_gradients, | ||
| 39 | + colocate_gradients_with_ops: colocate_gradients_with_ops); | ||
| 40 | + | ||
| 36 | 41 | return this; | |
| 37 | 42 | } | |
| 38 | 43 | ||
@@ -41,15 +46,30 @@ public Optimizer minimize(Tensor loss, GateGradientType gate_gradients = GateGra | |||
| 41 | 46 | /// </summary> | |
| 42 | 47 | /// <param name="loss"></param> | |
| 43 | 48 | /// <param name="gate_gradients"></param> | |
| 44 | - public List<KeyValuePair<object, object>> compute_gradients(Tensor loss, GateGradientType gate_gradients = GateGradientType.GATE_OP) | ||
| 49 | + public List<KeyValuePair<object, object>> compute_gradients(Tensor loss, | ||
| 50 | + List<RefVariable> var_list = null, | ||
| 51 | + GateGradientType gate_gradients = GateGradientType.GATE_OP, | ||
| 52 | + bool colocate_gradients_with_ops = false) | ||
| 45 | 53 | { | |
| 46 | 54 | int num_towers = 1; | |
| 47 | 55 | if(distribute_lib.get_loss_reduction() == VariableAggregationType.MEAN) | |
| 48 | 56 | { | |
| 49 | 57 | ||
| 50 | 58 | } | |
| 51 | 59 | ||
| 52 | - var var_list = variables.trainable_variables(); | ||
| 60 | + var tmp = variables.trainable_variables(); | ||
| 61 | + switch (tmp) | ||
| 62 | + { | ||
| 63 | + case List<RefVariable> values: | ||
| 64 | + var_list = values; | ||
| 65 | + break; | ||
| 66 | + } | ||
| 67 | + | ||
| 68 | + foreach(var v in var_list) | ||
| 69 | + { | ||
| 70 | + | ||
| 71 | + } | ||
| 72 | + | ||
| 53 | 73 | return null; | |
| 54 | 74 | } | |
| 55 | 75 | } | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -64,6 +64,28 @@ private void _init_from_args(object initial_value, | |||
| 64 | 64 | var shape = _initial_value.shape; | |
| 65 | 65 | dtype = _initial_value.dtype; | |
| 66 | 66 | _variable = gen_state_ops.variable_v2(shape, dtype, name); | |
| 67 | + | ||
| 68 | + // Manually overrides the variable's shape with the initial value's. | ||
| 69 | + if (validate_shape) | ||
| 70 | + { | ||
| 71 | + var initial_value_shape = _initial_value.shape; | ||
| 72 | + } | ||
| 73 | + | ||
| 74 | + // If 'initial_value' makes use of other variables, make sure we don't | ||
| 75 | + // have an issue if these other variables aren't initialized first by | ||
| 76 | + // using their initialized_value() method. | ||
| 77 | + | ||
| 78 | + ops.add_to_collections(collections, this); | ||
| 79 | + } | ||
| 80 | + | ||
| 81 | + public static implicit operator _VariableScopeStore(RefVariable variable) | ||
| 82 | + { | ||
| 83 | + return null; | ||
| 84 | + } | ||
| 85 | + | ||
| 86 | + public static implicit operator RefVariable(_VariableScopeStore store) | ||
| 87 | + { | ||
| 88 | + return null; | ||
| 67 | 89 | } | |
| 68 | 90 | } | |
| 69 | 91 | } | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -23,6 +23,8 @@ public static Tensor variable_v2(long[] shape, TF_DataType dtype, string name = | |||
| 23 | 23 | var keywords = new Dictionary<string, object>(); | |
| 24 | 24 | keywords.Add("dtype", dtype); | |
| 25 | 25 | keywords.Add("shape", shape); | |
| 26 | + keywords.Add("container", container); | ||
| 27 | + keywords.Add("shared_name", shared_name); | ||
| 26 | 28 | ||
| 27 | 29 | var _op = _op_def_lib._apply_op_helper("VariableV2", name: name, keywords: keywords); | |
| 28 | 30 | ||
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -39,18 +39,30 @@ public static VariableScope get_variable_scope() | |||
| 39 | 39 | ||
| 40 | 40 | public static _VariableScopeStore get_variable_scope_store() | |
| 41 | 41 | { | |
| 42 | + _VariableScopeStore ret = null; | ||
| 42 | 43 | var scope_store = ops.get_collection(_VARSCOPESTORE_KEY); | |
| 43 | 44 | if (scope_store == null) | |
| 44 | 45 | { | |
| 45 | - scope_store = new _VariableScopeStore(); | ||
| 46 | - ops.add_to_collection(_VARSCOPESTORE_KEY, scope_store); | ||
| 46 | + ret = new _VariableScopeStore(); | ||
| 47 | + ops.add_to_collection(_VARSCOPESTORE_KEY, ret); | ||
| 47 | 48 | } | |
| 48 | 49 | else | |
| 49 | 50 | { | |
| 50 | - // scope_store = scope_store[0]; | ||
| 51 | + switch (scope_store) | ||
| 52 | + { | ||
| 53 | + case List<RefVariable> values: | ||
| 54 | + ret = values[0]; | ||
| 55 | + break; | ||
| 56 | + case List<_VariableScopeStore> values: | ||
| 57 | + ret = values[0]; | ||
| 58 | + break; | ||
| 59 | + default: | ||
| 60 | + throw new InvalidOperationException("get_variable_scope_store"); | ||
| 61 | + } | ||
| 62 | + | ||
| 51 | 63 | } | |
| 52 | 64 | ||
| 53 | - return scope_store; | ||
| 65 | + return ret; | ||
| 54 | 66 | } | |
| 55 | 67 | ||
| 56 | 68 | public static bool _get_trainable_value(VariableSynchronization synchronization, bool? trainable = null) | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -14,7 +14,7 @@ public class name_scope | |||
| 14 | 14 | public Context _ctx; | |
| 15 | 15 | public string _name_scope; | |
| 16 | 16 | ||
| 17 | - public name_scope(string name, string default_name, List<object> values) | ||
| 17 | + public name_scope(string name, string default_name = "", List<object> values = null) | ||
| 18 | 18 | { | |
| 19 | 19 | _name = name; | |
| 20 | 20 | _default_name = default_name; | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -12,15 +12,21 @@ namespace Tensorflow | |||
| 12 | 12 | { | |
| 13 | 13 | public partial class ops | |
| 14 | 14 | { | |
| 15 | - public static void add_to_collection(string name, object value) | ||
| 15 | + public static void add_to_collection<T>(string name, T value) | ||
| 16 | 16 | { | |
| 17 | 17 | var graph = tf.get_default_graph(); | |
| 18 | 18 | graph.add_to_collection(name, value); | |
| 19 | 19 | } | |
| 20 | 20 | ||
| 21 | - public static _VariableScopeStore get_collection(string key) | ||
| 21 | + public static void add_to_collections<T>(List<string> names, T value) | ||
| 22 | 22 | { | |
| 23 | - return null;// get_default_graph().get_collection(key); | ||
| 23 | + var graph = tf.get_default_graph(); | ||
| 24 | + graph.add_to_collections(names, value); | ||
| 25 | + } | ||
| 26 | + | ||
| 27 | + public static object get_collection(string key) | ||
| 28 | + { | ||
| 29 | + return get_default_graph().get_collection(key); | ||
| 24 | 30 | } | |
| 25 | 31 | ||
| 26 | 32 | public static Graph get_default_graph() | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -27,12 +27,12 @@ public void Run() | |||
| 27 | 27 | var train_Y = np.array(1.7, 2.76, 2.09, 3.19, 1.694, 1.573, 3.366, 2.596, 2.53, 1.221, | |
| 28 | 28 | 2.827, 3.465, 1.65, 2.904, 2.42, 2.94, 1.3); | |
| 29 | 29 | var n_samples = train_X.shape[0]; | |
| 30 | - | ||
| 30 | + | ||
| 31 | 31 | // tf Graph Input | |
| 32 | 32 | var X = tf.placeholder(tf.float64); | |
| 33 | 33 | var Y = tf.placeholder(tf.float64); | |
| 34 | 34 | ||
| 35 | - // Set model weights | ||
| 35 | + // Set model weights | ||
| 36 | 36 | var W = tf.Variable(rng.randn<double>(), name: "weight"); | |
| 37 | 37 | var b = tf.Variable(rng.randn<double>(), name: "bias"); | |
| 38 | 38 | ||
| Back | FazBrowse Home | New Git URL |
0 commit comments