FazBrowse GitHub Viewer | Trending |
URL:
| Home
Tools: [Download Repo ZIP]   [Original HTTPS Page]

fix -1 index for _ConcatGradHelper · attackgithub/TensorFlow.NET@b357cad · GitHub

Repository navigation

Commit b357cad

Browse files
committed
fix -1 index for _ConcatGradHelper
1 parent 02f0f25 commit b357cad

9 files changed

Lines changed: 54 additions & 23 deletions

File tree

‎src/TensorFlowNET.Core/Gradients/array_grad.cs‎

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -40,6 +40,8 @@ private static Tensor[] _ConcatGradHelper(Operation op, Tensor grad, int start_v
4040
return end_value_index <= dim_index ? new Tensor[] { grad, null } : new Tensor[] { null, grad };
4141

4242
var concat_dim = op.inputs[dim_index];
43+
if (end_value_index == -1)
44+
end_value_index = op.inputs.Length - 1;
4345
var input_values = op.inputs._inputs.Skip(start_value_index).Take(end_value_index - start_value_index).ToArray();
4446

4547
var out_grads = new List<Tensor>();

‎src/TensorFlowNET.Core/Gradients/gradients_impl.py.cs‎

Lines changed: 3 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -43,12 +43,6 @@ public static Tensor[] _GradientsHelper(Tensor[] ys,
4343
if (grad_ys == null)
4444
grad_ys = new Tensor[ys.Length];
4545

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-
5246
// Iterate over the collected ops.
5347
/**
5448
* grads: op => list of gradients received on each output endpoint of the
@@ -59,7 +53,8 @@ public static Tensor[] _GradientsHelper(Tensor[] ys,
5953
**/
6054
var grads = new Dictionary<string, Tensor[][]>();
6155

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 =>
6358
{
6459
string grad_scope = scope;
6560
// Get a uid for this call to gradients that can be used to help
@@ -166,7 +161,7 @@ public static Tensor[] _GradientsHelper(Tensor[] ys,
166161
}
167162

168163
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))
170165
{
171166
if(in_grad != null)
172167
{

‎src/TensorFlowNET.Core/Graphs/Graph.cs‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -72,7 +72,7 @@ public partial class Graph : IPython, IDisposable
7272
private string _graph_key;
7373
public string graph_key => _graph_key;
7474
public string _last_loss_reduction;
75-
75+
public bool _is_loss_scaled_by_optimizer { get; set; }
7676
public Status Status { get; }
7777

7878
/// <summary>

‎src/TensorFlowNET.Core/Operations/InputList.cs‎

Lines changed: 9 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -10,7 +10,15 @@ public class InputList : IEnumerable
1010
{
1111
public Tensor[] _inputs;
1212
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+
}
1422

1523
public InputList(Tensor[] inputs)
1624
{

‎src/TensorFlowNET.Core/TensorFlowNET.Core.csproj‎

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -56,6 +56,7 @@ Removed global static graph instance.</PackageReleaseNotes>
5656
</ItemGroup>
5757

5858
<ItemGroup>
59+
<Folder Include="Distribute\" />
5960
<Folder Include="Keras\Initializers\" />
6061
</ItemGroup>
6162

‎src/TensorFlowNET.Core/Train/Optimizer.cs‎

Lines changed: 20 additions & 11 deletions
Original file line numberDiff line numberDiff line change
@@ -2,7 +2,6 @@
22
using System.Collections.Generic;
33
using System.Linq;
44
using System.Text;
5-
using distribute_lib = Tensorflow.Distribute;
65
using static Tensorflow.Python;
76

87
namespace Tensorflow
@@ -82,7 +81,8 @@ public Operation minimize(Tensor loss,
8281
var grads_and_vars = compute_gradients(loss, var_list:var_list,
8382
gate_gradients: gate_gradients,
8483
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);
8686

8787
var vars_with_grad = grads_and_vars.Where(x => x.Item1 != null).Select(x => x.Item2).ToArray();
8888
if (vars_with_grad.Length == 0)
@@ -232,30 +232,31 @@ public Tuple<Tensor, RefVariable>[] compute_gradients(Tensor loss,
232232
int? aggregation_method = null,
233233
GateGradientType gate_gradients = GateGradientType.GATE_OP,
234234
bool colocate_gradients_with_ops = false,
235-
Tensor[] grad_loss = null)
235+
Tensor grad_loss = null)
236236
{
237+
// Scale loss if using a "mean" loss reduction and multiple replicas.
238+
loss = _scale_loss(loss);
237239
int num_towers = 1;
238-
if(distribute_lib.get_loss_reduction() == VariableAggregationType.MEAN)
239-
{
240-
241-
}
240+
242241

243242
var tmp = variables.trainable_variables();
243+
var vars = ops.get_collection<RefVariable>(ops.GraphKeys.TRAINABLE_RESOURCE_VARIABLES);
244244
switch (tmp)
245245
{
246246
case List<RefVariable> values:
247-
var_list = values;
247+
var_list = values.Concat(vars).ToList();
248248
break;
249249
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();
251251
break;
252252
}
253253

254+
var_list = var_list.Concat(ops.get_collection<RefVariable>(ops.GraphKeys._STREAMING_MODEL_PORTS)).ToList();
254255
var processors = var_list.Select(v => optimizer._get_processor(v)).ToList();
255256
var var_refs = processors.Select(x => x.target()).ToArray();
256257

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,
259260
aggregation_method: aggregation_method,
260261
colocate_gradients_with_ops: colocate_gradients_with_ops);
261262

@@ -269,6 +270,14 @@ public Tuple<Tensor, RefVariable>[] compute_gradients(Tensor loss,
269270
return grads_and_vars;
270271
}
271272

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+
272281
protected T _call_if_callable<T>(T param)
273282
{
274283
return param;

‎src/TensorFlowNET.Core/ops.GraphKeys.cs‎

Lines changed: 10 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -22,6 +22,16 @@ public static class GraphKeys
2222
/// </summary>
2323
public static string TRAINABLE_VARIABLES = "trainable_variables";
2424

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+
2535
/// <summary>
2636
/// Key to collect losses
2737
/// </summary>

‎src/TensorFlowNET.Core/ops.py.cs‎

Lines changed: 5 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -45,6 +45,11 @@ public static object get_collection(string key, string scope = null)
4545
return get_default_graph().get_collection(key, scope);
4646
}
4747

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+
4853
public static object get_collection_ref(string key)
4954
{
5055
return get_default_graph().get_collection_ref(key);

‎test/TensorFlowNET.Examples/TextProcess/CnnTextClassification.cs‎

Lines changed: 3 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -207,15 +207,16 @@ public Graph BuildGraph()
207207
Tensor predictions = null;
208208
with(tf.name_scope("output"), delegate
209209
{
210-
logits = tf.layers.dense(h_pool_flat, keep_prob);
210+
logits = tf.layers.dense(h_pool_flat, NUM_CLASS);
211211
predictions = tf.argmax(logits, -1, output_type: tf.int32);
212212
});
213213

214214
with(tf.name_scope("loss"), delegate
215215
{
216216
var sscel = tf.nn.sparse_softmax_cross_entropy_with_logits(logits: logits, labels: y);
217217
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);
219220
});
220221

221222
with(tf.name_scope("accuracy"), delegate

0 commit comments

Comments
 (0)

Back | FazBrowse Home | New Git URL