| FazBrowse GitHub Viewer | Trending | | Home |
| Tools: [Download Repo ZIP] [Original HTTPS Page] |
1 parent af73e3c commit ecbda0c
8 files changed
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -31,7 +31,7 @@ public Optimizer GradientDescentOptimizer(float learning_rate) | |||
| 31 | 31 | public Optimizer AdamOptimizer(float learning_rate, string name = "Adam") | |
| 32 | 32 | => new AdamOptimizer(learning_rate, name: name); | |
| 33 | 33 | ||
| 34 | - public object ExponentialMovingAverage(float decay) | ||
| 34 | + public ExponentialMovingAverage ExponentialMovingAverage(float decay) | ||
| 35 | 35 | => new ExponentialMovingAverage(decay); | |
| 36 | 36 | ||
| 37 | 37 | public Saver Saver(VariableV1[] var_list = null) => new Saver(var_list: var_list); | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -1,6 +1,8 @@ | |||
| 1 | 1 | using System; | |
| 2 | 2 | using System.Collections.Generic; | |
| 3 | + using System.Linq; | ||
| 3 | 4 | using System.Text; | |
| 5 | + using static Tensorflow.Binding; | ||
| 4 | 6 | ||
| 5 | 7 | namespace Tensorflow.Train | |
| 6 | 8 | { | |
@@ -11,6 +13,7 @@ public class ExponentialMovingAverage | |||
| 11 | 13 | bool _zero_debias; | |
| 12 | 14 | string _name; | |
| 13 | 15 | public string name => _name; | |
| 16 | + List<VariableV1> _averages; | ||
| 14 | 17 | ||
| 15 | 18 | public ExponentialMovingAverage(float decay, int? num_updates = null, bool zero_debias = false, | |
| 16 | 19 | string name = "ExponentialMovingAverage") | |
@@ -19,18 +22,31 @@ public ExponentialMovingAverage(float decay, int? num_updates = null, bool zero_ | |||
| 19 | 22 | _num_updates = num_updates; | |
| 20 | 23 | _zero_debias = zero_debias; | |
| 21 | 24 | _name = name; | |
| 25 | + _averages = new List<VariableV1>(); | ||
| 22 | 26 | } | |
| 23 | 27 | ||
| 24 | 28 | /// <summary> | |
| 25 | 29 | /// Maintains moving averages of variables. | |
| 26 | 30 | /// </summary> | |
| 27 | 31 | /// <param name="var_list"></param> | |
| 28 | 32 | /// <returns></returns> | |
| 29 | - public Operation apply(VariableV1[] var_list = null) | ||
| 33 | + public Operation apply(RefVariable[] var_list = null) | ||
| 30 | 34 | { | |
| 31 | - throw new NotImplementedException(""); | ||
| 32 | - } | ||
| 35 | + if (var_list == null) | ||
| 36 | + var_list = variables.trainable_variables() as RefVariable[]; | ||
| 33 | 37 | ||
| 38 | + foreach(var var in var_list) | ||
| 39 | + { | ||
| 40 | + if (!_averages.Contains(var)) | ||
| 41 | + { | ||
| 42 | + ops.init_scope(); | ||
| 43 | + var slot = new SlotCreator(); | ||
| 44 | + var.initialized_value(); | ||
| 45 | + // var avg = slot.create_zeros_slot | ||
| 46 | + } | ||
| 47 | + } | ||
| 34 | 48 | ||
| 49 | + throw new NotImplementedException(""); | ||
| 50 | + } | ||
| 35 | 51 | } | |
| 36 | 52 | } | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -308,5 +308,28 @@ public RefVariable from_proto(VariableDef proto, string import_scope) | |||
| 308 | 308 | { | |
| 309 | 309 | throw new NotImplementedException(); | |
| 310 | 310 | } | |
| 311 | + | ||
| 312 | + /// <summary> | ||
| 313 | + /// Returns the value of this variable, read in the current context. | ||
| 314 | + /// </summary> | ||
| 315 | + /// <returns></returns> | ||
| 316 | + private ITensorOrOperation read_value() | ||
| 317 | + { | ||
| 318 | + return array_ops.identity(_variable, name: "read"); | ||
| 319 | + } | ||
| 320 | + | ||
| 321 | + public Tensor is_variable_initialized(RefVariable variable) | ||
| 322 | + { | ||
| 323 | + return state_ops.is_variable_initialized(variable); | ||
| 324 | + } | ||
| 325 | + | ||
| 326 | + public Tensor initialized_value() | ||
| 327 | + { | ||
| 328 | + ops.init_scope(); | ||
| 329 | + throw new NotImplementedException(""); | ||
| 330 | + /*return control_flow_ops.cond(is_variable_initialized(this), | ||
| 331 | + read_value, | ||
| 332 | + () => initial_value);*/ | ||
| 333 | + } | ||
| 311 | 334 | } | |
| 312 | 335 | } | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -14,6 +14,7 @@ You may obtain a copy of the License at | |||
| 14 | 14 | limitations under the License. | |
| 15 | 15 | ******************************************************************************/ | |
| 16 | 16 | ||
| 17 | + using System; | ||
| 17 | 18 | using System.Collections.Generic; | |
| 18 | 19 | using Tensorflow.Eager; | |
| 19 | 20 | ||
@@ -145,5 +146,10 @@ public static Tensor scatter_add(RefVariable @ref, Tensor indices, Tensor update | |||
| 145 | 146 | var _op = _op_def_lib._apply_op_helper("ScatterAdd", name: name, args: new { @ref, indices, updates, use_locking }); | |
| 146 | 147 | return _op.outputs[0]; | |
| 147 | 148 | } | |
| 149 | + | ||
| 150 | + public static Tensor is_variable_initialized(RefVariable @ref, string name = null) | ||
| 151 | + { | ||
| 152 | + throw new NotImplementedException(""); | ||
| 153 | + } | ||
| 148 | 154 | } | |
| 149 | 155 | } | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -106,5 +106,13 @@ public static Tensor scatter_add(RefVariable @ref, Tensor indices, Tensor update | |||
| 106 | 106 | ||
| 107 | 107 | throw new NotImplementedException("scatter_add"); | |
| 108 | 108 | } | |
| 109 | + | ||
| 110 | + public static Tensor is_variable_initialized(RefVariable @ref, string name = null) | ||
| 111 | + { | ||
| 112 | + if (@ref.dtype.is_ref_dtype()) | ||
| 113 | + return gen_state_ops.is_variable_initialized(@ref: @ref, name: name); | ||
| 114 | + throw new NotImplementedException(""); | ||
| 115 | + //return @ref.is_initialized(name: name); | ||
| 116 | + } | ||
| 109 | 117 | } | |
| 110 | 118 | } | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -91,7 +91,12 @@ public Graph BuildGraph() | |||
| 91 | 91 | ||
| 92 | 92 | tf_with(tf.name_scope("define_loss"), scope => | |
| 93 | 93 | { | |
| 94 | - model = new YOLOv3(cfg, input_data, trainable); | ||
| 94 | + // model = new YOLOv3(cfg, input_data, trainable); | ||
| 95 | + }); | ||
| 96 | + | ||
| 97 | + tf_with(tf.name_scope("define_weight_decay"), scope => | ||
| 98 | + { | ||
| 99 | + var moving_ave = tf.train.ExponentialMovingAverage(moving_ave_decay).apply((RefVariable[])tf.trainable_variables()); | ||
| 95 | 100 | }); | |
| 96 | 101 | ||
| 97 | 102 | return graph; | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -8,6 +8,9 @@ | |||
| 8 | 8 | ||
| 9 | 9 | namespace TensorFlowNET.UnitTest | |
| 10 | 10 | { | |
| 11 | + /// <summary> | ||
| 12 | + /// Find more examples in https://www.programcreek.com/python/example/90444/tensorflow.read_file | ||
| 13 | + /// </summary> | ||
| 11 | 14 | [TestClass] | |
| 12 | 15 | public class ImageTest | |
| 13 | 16 | { | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -69,9 +69,7 @@ public void NestedNameScope_Using() | |||
| 69 | 69 | Assert.AreEqual("scope1", g._name_stack); | |
| 70 | 70 | var const3 = tf.constant(2.0); | |
| 71 | 71 | Assert.AreEqual("scope1/Const_1:0", const3.name); | |
| 72 | - } | ||
| 73 | - | ||
| 74 | - ; | ||
| 72 | + }; | ||
| 75 | 73 | ||
| 76 | 74 | g.Dispose(); | |
| 77 | 75 | ||
| Back | FazBrowse Home | New Git URL |
0 commit comments