| FazBrowse GitHub Viewer | Trending | | Home |
| Tools: [Download Repo ZIP] [Original HTTPS Page] |
1 parent 02f0f25 commit b357cad
9 files changed
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -40,6 +40,8 @@ private static Tensor[] _ConcatGradHelper(Operation op, Tensor grad, int start_v | |||
| 40 | 40 | return end_value_index <= dim_index ? new Tensor[] { grad, null } : new Tensor[] { null, grad }; | |
| 41 | 41 | ||
| 42 | 42 | var concat_dim = op.inputs[dim_index]; | |
| 43 | + if (end_value_index == -1) | ||
| 44 | + end_value_index = op.inputs.Length - 1; | ||
| 43 | 45 | var input_values = op.inputs._inputs.Skip(start_value_index).Take(end_value_index - start_value_index).ToArray(); | |
| 44 | 46 | ||
| 45 | 47 | var out_grads = new List<Tensor>(); | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -43,12 +43,6 @@ public static Tensor[] _GradientsHelper(Tensor[] ys, | |||
| 43 | 43 | if (grad_ys == null) | |
| 44 | 44 | grad_ys = new Tensor[ys.Length]; | |
| 45 | 45 | ||
| 46 | - var all = new List<Tensor>(); | ||
| 47 | - all.AddRange(ys); | ||
| 48 | - all.AddRange(xs); | ||
| 49 | - all.AddRange(stop_gradients); | ||
| 50 | - all.AddRange(grad_ys); | ||
| 51 | - | ||
| 52 | 46 | // Iterate over the collected ops. | |
| 53 | 47 | /** | |
| 54 | 48 | * grads: op => list of gradients received on each output endpoint of the | |
@@ -59,7 +53,8 @@ public static Tensor[] _GradientsHelper(Tensor[] ys, | |||
| 59 | 53 | **/ | |
| 60 | 54 | var grads = new Dictionary<string, Tensor[][]>(); | |
| 61 | 55 | ||
| 62 | - with(ops.name_scope(name, "gradients", values: all), scope => | ||
| 56 | + with(ops.name_scope(name, "gradients", | ||
| 57 | + values: ys.Concat(xs).Concat(stop_gradients).Concat(grad_ys)), scope => | ||
| 63 | 58 | { | |
| 64 | 59 | string grad_scope = scope; | |
| 65 | 60 | // Get a uid for this call to gradients that can be used to help | |
@@ -166,7 +161,7 @@ public static Tensor[] _GradientsHelper(Tensor[] ys, | |||
| 166 | 161 | } | |
| 167 | 162 | ||
| 168 | 163 | var inputs = _NonEagerInputs(op, xs).ToList(); | |
| 169 | - foreach (var (t_in, in_grad) in Python.zip(inputs, in_grads)) | ||
| 164 | + foreach (var (t_in, in_grad) in zip(inputs, in_grads)) | ||
| 170 | 165 | { | |
| 171 | 166 | if(in_grad != null) | |
| 172 | 167 | { | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -72,7 +72,7 @@ public partial class Graph : IPython, IDisposable | |||
| 72 | 72 | private string _graph_key; | |
| 73 | 73 | public string graph_key => _graph_key; | |
| 74 | 74 | public string _last_loss_reduction; | |
| 75 | - | ||
| 75 | + public bool _is_loss_scaled_by_optimizer { get; set; } | ||
| 76 | 76 | public Status Status { get; } | |
| 77 | 77 | ||
| 78 | 78 | /// <summary> | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -10,7 +10,15 @@ public class InputList : IEnumerable | |||
| 10 | 10 | { | |
| 11 | 11 | public Tensor[] _inputs; | |
| 12 | 12 | public int Length => _inputs.Length; | |
| 13 | - public Tensor this[int index] => _inputs[index]; | ||
| 13 | + public Tensor this[int index] | ||
| 14 | + { | ||
| 15 | + get | ||
| 16 | + { | ||
| 17 | + if (index == -1) | ||
| 18 | + index = _inputs.Length - 1; | ||
| 19 | + return _inputs[index]; | ||
| 20 | + } | ||
| 21 | + } | ||
| 14 | 22 | ||
| 15 | 23 | public InputList(Tensor[] inputs) | |
| 16 | 24 | { | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -56,6 +56,7 @@ Removed global static graph instance.</PackageReleaseNotes> | |||
| 56 | 56 | </ItemGroup> | |
| 57 | 57 | ||
| 58 | 58 | <ItemGroup> | |
| 59 | + <Folder Include="Distribute\" /> | ||
| 59 | 60 | <Folder Include="Keras\Initializers\" /> | |
| 60 | 61 | </ItemGroup> | |
| 61 | 62 | ||
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -2,7 +2,6 @@ | |||
| 2 | 2 | using System.Collections.Generic; | |
| 3 | 3 | using System.Linq; | |
| 4 | 4 | using System.Text; | |
| 5 | - using distribute_lib = Tensorflow.Distribute; | ||
| 6 | 5 | using static Tensorflow.Python; | |
| 7 | 6 | ||
| 8 | 7 | namespace Tensorflow | |
@@ -82,7 +81,8 @@ public Operation minimize(Tensor loss, | |||
| 82 | 81 | var grads_and_vars = compute_gradients(loss, var_list:var_list, | |
| 83 | 82 | gate_gradients: gate_gradients, | |
| 84 | 83 | aggregation_method:aggregation_method, | |
| 85 | - colocate_gradients_with_ops: colocate_gradients_with_ops); | ||
| 84 | + colocate_gradients_with_ops: colocate_gradients_with_ops, | ||
| 85 | + grad_loss: grad_loss); | ||
| 86 | 86 | ||
| 87 | 87 | var vars_with_grad = grads_and_vars.Where(x => x.Item1 != null).Select(x => x.Item2).ToArray(); | |
| 88 | 88 | if (vars_with_grad.Length == 0) | |
@@ -232,30 +232,31 @@ public Tuple<Tensor, RefVariable>[] compute_gradients(Tensor loss, | |||
| 232 | 232 | int? aggregation_method = null, | |
| 233 | 233 | GateGradientType gate_gradients = GateGradientType.GATE_OP, | |
| 234 | 234 | bool colocate_gradients_with_ops = false, | |
| 235 | - Tensor[] grad_loss = null) | ||
| 235 | + Tensor grad_loss = null) | ||
| 236 | 236 | { | |
| 237 | + // Scale loss if using a "mean" loss reduction and multiple replicas. | ||
| 238 | + loss = _scale_loss(loss); | ||
| 237 | 239 | int num_towers = 1; | |
| 238 | - if(distribute_lib.get_loss_reduction() == VariableAggregationType.MEAN) | ||
| 239 | - { | ||
| 240 | - | ||
| 241 | - } | ||
| 240 | + | ||
| 242 | 241 | ||
| 243 | 242 | var tmp = variables.trainable_variables(); | |
| 243 | + var vars = ops.get_collection<RefVariable>(ops.GraphKeys.TRAINABLE_RESOURCE_VARIABLES); | ||
| 244 | 244 | switch (tmp) | |
| 245 | 245 | { | |
| 246 | 246 | case List<RefVariable> values: | |
| 247 | - var_list = values; | ||
| 247 | + var_list = values.Concat(vars).ToList(); | ||
| 248 | 248 | break; | |
| 249 | 249 | case List<VariableV1> values: | |
| 250 | - var_list = values.Select(x => x as RefVariable).ToList(); | ||
| 250 | + var_list = values.Select(x => x as RefVariable).Concat(vars).ToList(); | ||
| 251 | 251 | break; | |
| 252 | 252 | } | |
| 253 | 253 | ||
| 254 | + var_list = var_list.Concat(ops.get_collection<RefVariable>(ops.GraphKeys._STREAMING_MODEL_PORTS)).ToList(); | ||
| 254 | 255 | var processors = var_list.Select(v => optimizer._get_processor(v)).ToList(); | |
| 255 | 256 | var var_refs = processors.Select(x => x.target()).ToArray(); | |
| 256 | 257 | ||
| 257 | - var grads = gradients_impl.gradients(new Tensor[] { loss }, var_refs, grad_ys: grad_loss, | ||
| 258 | - gate_gradients: (gate_gradients == GateGradientType.GATE_OP), | ||
| 258 | + var grads = gradients_impl.gradients(new Tensor[] { loss }, var_refs, grad_ys: grad_loss == null ? null : new Tensor[] { grad_loss }, | ||
| 259 | + gate_gradients: gate_gradients == GateGradientType.GATE_OP, | ||
| 259 | 260 | aggregation_method: aggregation_method, | |
| 260 | 261 | colocate_gradients_with_ops: colocate_gradients_with_ops); | |
| 261 | 262 | ||
@@ -269,6 +270,14 @@ public Tuple<Tensor, RefVariable>[] compute_gradients(Tensor loss, | |||
| 269 | 270 | return grads_and_vars; | |
| 270 | 271 | } | |
| 271 | 272 | ||
| 273 | + private Tensor _scale_loss(Tensor loss_value) | ||
| 274 | + { | ||
| 275 | + ops.get_default_graph()._is_loss_scaled_by_optimizer = false; | ||
| 276 | + // TODO | ||
| 277 | + // if distribute_lib.get_loss_reduction() == ds_reduce_util.ReduceOp.MEAN: | ||
| 278 | + return loss_value; | ||
| 279 | + } | ||
| 280 | + | ||
| 272 | 281 | protected T _call_if_callable<T>(T param) | |
| 273 | 282 | { | |
| 274 | 283 | return param; | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -22,6 +22,16 @@ public static class GraphKeys | |||
| 22 | 22 | /// </summary> | |
| 23 | 23 | public static string TRAINABLE_VARIABLES = "trainable_variables"; | |
| 24 | 24 | ||
| 25 | + /// <summary> | ||
| 26 | + /// Trainable resource-style variables. | ||
| 27 | + /// </summary> | ||
| 28 | + public static string TRAINABLE_RESOURCE_VARIABLES = "trainable_resource_variables"; | ||
| 29 | + | ||
| 30 | + /// <summary> | ||
| 31 | + /// Key for streaming model ports. | ||
| 32 | + /// </summary> | ||
| 33 | + public static string _STREAMING_MODEL_PORTS = "streaming_model_ports"; | ||
| 34 | + | ||
| 25 | 35 | /// <summary> | |
| 26 | 36 | /// Key to collect losses | |
| 27 | 37 | /// </summary> | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -45,6 +45,11 @@ public static object get_collection(string key, string scope = null) | |||
| 45 | 45 | return get_default_graph().get_collection(key, scope); | |
| 46 | 46 | } | |
| 47 | 47 | ||
| 48 | + public static List<T> get_collection<T>(string key, string scope = null) | ||
| 49 | + { | ||
| 50 | + return get_default_graph().get_collection<T>(key, scope); | ||
| 51 | + } | ||
| 52 | + | ||
| 48 | 53 | public static object get_collection_ref(string key) | |
| 49 | 54 | { | |
| 50 | 55 | return get_default_graph().get_collection_ref(key); | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -207,15 +207,16 @@ public Graph BuildGraph() | |||
| 207 | 207 | Tensor predictions = null; | |
| 208 | 208 | with(tf.name_scope("output"), delegate | |
| 209 | 209 | { | |
| 210 | - logits = tf.layers.dense(h_pool_flat, keep_prob); | ||
| 210 | + logits = tf.layers.dense(h_pool_flat, NUM_CLASS); | ||
| 211 | 211 | predictions = tf.argmax(logits, -1, output_type: tf.int32); | |
| 212 | 212 | }); | |
| 213 | 213 | ||
| 214 | 214 | with(tf.name_scope("loss"), delegate | |
| 215 | 215 | { | |
| 216 | 216 | var sscel = tf.nn.sparse_softmax_cross_entropy_with_logits(logits: logits, labels: y); | |
| 217 | 217 | var loss = tf.reduce_mean(sscel); | |
| 218 | - var optimizer = tf.train.AdamOptimizer(learning_rate).minimize(loss, global_step: global_step); | ||
| 218 | + var adam = tf.train.AdamOptimizer(learning_rate); | ||
| 219 | + var optimizer = adam.minimize(loss, global_step: global_step); | ||
| 219 | 220 | }); | |
| 220 | 221 | ||
| 221 | 222 | with(tf.name_scope("accuracy"), delegate | |
| Back | FazBrowse Home | New Git URL |
0 commit comments