| FazBrowse GitHub Viewer | Trending | | Home |
| Tools: [Download Repo ZIP] [Original HTTPS Page] |
1 parent 93fd34b commit e19e59b
10 files changed
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -35,6 +35,7 @@ public sealed class Context : IDisposable | |||
| 35 | 35 | public string ScopeName { get; set; } = ""; | |
| 36 | 36 | bool initialized = false; | |
| 37 | 37 | ContextSwitchStack context_switches; | |
| 38 | + public FunctionCallOptions FunctionCallOptions { get; } | ||
| 38 | 39 | ||
| 39 | 40 | public SafeContextHandle Handle { get; } | |
| 40 | 41 | ||
@@ -44,6 +45,7 @@ public Context(ContextOptions opts, Status status) | |||
| 44 | 45 | status.Check(true); | |
| 45 | 46 | context_switches = new ContextSwitchStack(defaultExecutionMode == EAGER_MODE, false); | |
| 46 | 47 | initialized = true; | |
| 48 | + FunctionCallOptions = new FunctionCallOptions(); | ||
| 47 | 49 | } | |
| 48 | 50 | ||
| 49 | 51 | /// <summary> | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -0,0 +1,20 @@ | |||
| 1 | + using System; | ||
| 2 | + using System.Collections.Generic; | ||
| 3 | + using System.Text; | ||
| 4 | + using Google.Protobuf; | ||
| 5 | + using Google.Protobuf.Collections; | ||
| 6 | + | ||
| 7 | + namespace Tensorflow.Contexts | ||
| 8 | + { | ||
| 9 | + public class FunctionCallOptions | ||
| 10 | + { | ||
| 11 | + public string config_proto_serialized() | ||
| 12 | + { | ||
| 13 | + var config = new ConfigProto | ||
| 14 | + { | ||
| 15 | + AllowSoftPlacement = true, | ||
| 16 | + }; | ||
| 17 | + return config.ToByteString().ToStringUtf8(); | ||
| 18 | + } | ||
| 19 | + } | ||
| 20 | + } | ||
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -371,7 +371,7 @@ bool SetOpAttrScalar(Context ctx, SafeOpHandle op, | |||
| 371 | 371 | switch (type) | |
| 372 | 372 | { | |
| 373 | 373 | case TF_AttrType.TF_ATTR_STRING: | |
| 374 | - c_api.TFE_OpSetAttrString(op, key, value.ToString(), (uint)value.ToString().Length); | ||
| 374 | + c_api.TFE_OpSetAttrString(op, key, value.ToString(), (ulong)value.ToString().Length); | ||
| 375 | 375 | break; | |
| 376 | 376 | case TF_AttrType.TF_ATTR_TYPE: | |
| 377 | 377 | c_api.TFE_OpSetAttrType(op, key, (TF_DataType)value); | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -241,7 +241,7 @@ public static void TFE_Execute(SafeOpHandle op, SafeTensorHandleHandle[] retvals | |||
| 241 | 241 | /// <param name="value">const void*</param> | |
| 242 | 242 | /// <param name="length">size_t</param> | |
| 243 | 243 | [DllImport(TensorFlowLibName)] | |
| 244 | - public static extern void TFE_OpSetAttrString(SafeOpHandle op, string attr_name, string value, uint length); | ||
| 244 | + public static extern void TFE_OpSetAttrString(SafeOpHandle op, string attr_name, string value, ulong length); | ||
| 245 | 245 | ||
| 246 | 246 | [DllImport(TensorFlowLibName)] | |
| 247 | 247 | public static extern void TFE_OpSetAttrTypeList(SafeOpHandle op, string attr_name, TF_DataType[] values, int num_values); | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -14,6 +14,7 @@ public class ConcreteFunction : IDisposable | |||
| 14 | 14 | { | |
| 15 | 15 | IntPtr _handle; | |
| 16 | 16 | FuncGraph func_graph; | |
| 17 | + public Tensor[] CapturedInputs => func_graph.external_captures(); | ||
| 17 | 18 | ||
| 18 | 19 | public string Name | |
| 19 | 20 | { | |
@@ -38,6 +39,8 @@ public ConcreteFunction(string name) | |||
| 38 | 39 | public ConcreteFunction(FuncGraph graph, Dictionary<string, string> attrs) | |
| 39 | 40 | { | |
| 40 | 41 | func_graph = graph; | |
| 42 | + | ||
| 43 | + ToGraph(graph.Inputs, graph.Outputs); | ||
| 41 | 44 | } | |
| 42 | 45 | ||
| 43 | 46 | public ConcreteFunction(Func<Tensor, Tensor> func, TF_DataType dtype) | |
@@ -124,6 +127,21 @@ public Tensors Invoke(Tensors inputs) | |||
| 124 | 127 | return flat_outputs; | |
| 125 | 128 | } | |
| 126 | 129 | ||
| 130 | + public Tensor[] CallFlat(Tensor[] args, Tensor[] captured_inputs) | ||
| 131 | + { | ||
| 132 | + var new_args = new List<Tensor>(); | ||
| 133 | + new_args.AddRange(args); | ||
| 134 | + new_args.AddRange(captured_inputs); | ||
| 135 | + args = new_args.ToArray(); | ||
| 136 | + | ||
| 137 | + var attrs = new object[] | ||
| 138 | + { | ||
| 139 | + "executor_type", "", | ||
| 140 | + "config_proto", tf.Context.FunctionCallOptions.config_proto_serialized() | ||
| 141 | + }; | ||
| 142 | + return tf.Runner.Execute(tf.Context, func_graph.FuncName, 1, args, attrs); | ||
| 143 | + } | ||
| 144 | + | ||
| 127 | 145 | ForwardBackwardCall SelectForwardAndBackwardFunctions(Tensors args, int possible_gradient_type, bool executing_eagerly) | |
| 128 | 146 | { | |
| 129 | 147 | var functions = new FirstOrderTapeGradientFunctions(func_graph, false); | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -48,7 +48,7 @@ public void Record(Tensors flat_outputs, Tensors inference_args) | |||
| 48 | 48 | getBackwardFunction: () => backward_function); | |
| 49 | 49 | } | |
| 50 | 50 | ||
| 51 | - (BackwardFunction, Tensors) _wrap_backward_function(FuncGraph forward_graph, ConcreteFunction backward, Tensors flat_outputs) | ||
| 51 | + (BackwardFunction, Tensors) _wrap_backward_function(FuncGraph forward_graph, ConcreteFunction backward, Tensors outputs) | ||
| 52 | 52 | { | |
| 53 | 53 | BackwardFunction _backward_function_wrapper = (output_grads, unneeded_gradients) => | |
| 54 | 54 | { | |
@@ -61,10 +61,11 @@ public void Record(Tensors flat_outputs, Tensors inference_args) | |||
| 61 | 61 | processed_args.add(arg); | |
| 62 | 62 | input_index += 1; | |
| 63 | 63 | } | |
| 64 | - return output_grads;// backward.Invoke(processed_args.ToArray()); | ||
| 64 | + | ||
| 65 | + return backward.CallFlat(processed_args.ToArray(), outputs); | ||
| 65 | 66 | }; | |
| 66 | 67 | ||
| 67 | - return (_backward_function_wrapper, flat_outputs); | ||
| 68 | + return (_backward_function_wrapper, outputs); | ||
| 68 | 69 | } | |
| 69 | 70 | ||
| 70 | 71 | protected (EagerDefinedFunction, FuncGraph, ConcreteFunction, List<int>, int) | |
@@ -82,24 +83,23 @@ public void Record(Tensors flat_outputs, Tensors inference_args) | |||
| 82 | 83 | } | |
| 83 | 84 | ||
| 84 | 85 | var gradients_wrt_outputs = new List<Tensor>(); | |
| 85 | - var backwards_graph = new FuncGraph($"{_BACKWARD_PREFIX}{_func_graph.FuncName}_{ops.uid()}"); | ||
| 86 | + var backwards_graph = new FuncGraph($"{_BACKWARD_PREFIX}_{ops.uid()}"); | ||
| 86 | 87 | foreach (var output in trainable_outputs) | |
| 87 | 88 | gradients_wrt_outputs.Add(tf.placeholder(output.dtype, output.shape)); | |
| 88 | 89 | var gradients_wrt_inputs = gradients_util._GradientsHelper(trainable_outputs.ToArray(), | |
| 89 | 90 | _func_graph.Inputs, | |
| 90 | 91 | grad_ys: gradients_wrt_outputs.ToArray(), | |
| 91 | 92 | src_graph: _func_graph); | |
| 92 | 93 | ||
| 93 | - tf.Context.restore_mode(); | ||
| 94 | - | ||
| 95 | - var forward_function_name = $"{_FORWARD_PREFIX}{_func_graph.FuncName}_{ops.uid()}"; | ||
| 94 | + var forward_function_name = $"{_FORWARD_PREFIX}_{ops.uid()}"; | ||
| 96 | 95 | var backward_function_attr = new Dictionary<string, string>(); | |
| 97 | 96 | backward_function_attr[FORWARD_FUNCTION_ATTRIBUTE_NAME] = forward_function_name; | |
| 97 | + gradients_wrt_outputs.append(backwards_graph.internal_captures()); | ||
| 98 | 98 | backwards_graph.Inputs = gradients_wrt_outputs; | |
| 99 | 99 | backwards_graph.Outputs = gradients_wrt_inputs; | |
| 100 | 100 | ||
| 101 | 101 | var backward_function = new ConcreteFunction(backwards_graph, backward_function_attr); | |
| 102 | - | ||
| 102 | + | ||
| 103 | 103 | var forward_function_attr = new Dictionary<string, string>(); | |
| 104 | 104 | forward_function_attr[BACKWARD_FUNCTION_ATTRIBUTE_NAME] = backward_function.Name; | |
| 105 | 105 | var forward_function = new EagerDefinedFunction(forward_function_name, _func_graph, | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -49,14 +49,14 @@ public static void RegisterFromAssembly() | |||
| 49 | 49 | RegisterGradientFunction(m.GetCustomAttribute<RegisterGradient>().Name, | |
| 50 | 50 | (oper, out_grads) => | |
| 51 | 51 | { | |
| 52 | - tf.Logger.Debug($"Caculate Gradient: {m.Name}"); | ||
| 52 | + tf.Logger.Debug($"Caculate Gradient: {oper.name} {m.Name}"); | ||
| 53 | 53 | var results = g.InvokeMember(m.Name, | |
| 54 | 54 | BindingFlags.InvokeMethod, | |
| 55 | 55 | null, | |
| 56 | 56 | null, | |
| 57 | 57 | args: new object[] { oper, out_grads }) as Tensor[]; | |
| 58 | 58 | foreach (var result in results.Where(x => x != null)) | |
| 59 | - tf.Logger.Debug($"{result.TensorShape}"); | ||
| 59 | + tf.Logger.Debug($"Gradient: {result.name} {result.TensorShape}"); | ||
| 60 | 60 | return results; | |
| 61 | 61 | } | |
| 62 | 62 | ); | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -26,7 +26,9 @@ public class FuncGraph : Graph | |||
| 26 | 26 | public Tensors Outputs { get; set; } | |
| 27 | 27 | public Dictionary<string, string> Attrs { get; set; } | |
| 28 | 28 | ||
| 29 | - Dictionary<long, (Tensor, Tensor)> _captures = new Dictionary<long, (Tensor, Tensor)>(); | ||
| 29 | + // new Dictionary<long, (Tensor, Tensor)> _captures = new Dictionary<long, (Tensor, Tensor)>(); | ||
| 30 | + // public new Tensor[] external_captures => _captures.Values.Select(x => x.Item1).ToArray(); | ||
| 31 | + | ||
| 30 | 32 | /// <summary> | |
| 31 | 33 | /// Construct a new FuncGraph. | |
| 32 | 34 | /// </summary> | |
@@ -129,7 +131,7 @@ Tensor capture(Tensor tensor, string name = null, TF_DataType shape = TF_DataTyp | |||
| 129 | 131 | Tensor _capture_helper(Tensor tensor, string name, TensorShape shape = null) | |
| 130 | 132 | { | |
| 131 | 133 | Tensor placeholder = null; | |
| 132 | - if (!_captures.ContainsKey(tensor.Id)) | ||
| 134 | + if (!_captures.Contains(tensor.Id)) | ||
| 133 | 135 | { | |
| 134 | 136 | placeholder = _create_substitute_placeholder(tensor, | |
| 135 | 137 | name: name, | |
@@ -139,7 +141,7 @@ Tensor _capture_helper(Tensor tensor, string name, TensorShape shape = null) | |||
| 139 | 141 | } | |
| 140 | 142 | else | |
| 141 | 143 | { | |
| 142 | - placeholder = _captures[tensor.Id].Item1; | ||
| 144 | + placeholder = (((Tensor, Tensor))_captures[tensor.Id]).Item2; | ||
| 143 | 145 | } | |
| 144 | 146 | ||
| 145 | 147 | BackwardFunction _backward_function_wrapper = (output_grads, unneeded_gradients) => | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -557,16 +557,16 @@ public static implicit operator IntPtr(Graph graph) | |||
| 557 | 557 | ||
| 558 | 558 | public Tensor[] external_captures() | |
| 559 | 559 | { | |
| 560 | - Tensor[] captures = new Tensor[this._captures.Count]; | ||
| 561 | - ICollection inner = this._captures.Keys; // c[0] | ||
| 560 | + Tensor[] captures = new Tensor[_captures.Count]; | ||
| 561 | + ICollection inner = _captures.Keys; // c[0] | ||
| 562 | 562 | inner.CopyTo(captures, 0); | |
| 563 | 563 | return captures; | |
| 564 | 564 | } | |
| 565 | 565 | ||
| 566 | 566 | public Tensor[] internal_captures() | |
| 567 | 567 | { | |
| 568 | - Tensor[] captures = new Tensor[this._captures.Count]; | ||
| 569 | - ICollection inner = this._captures.Values; // c[1] | ||
| 568 | + Tensor[] captures = new Tensor[_captures.Count]; | ||
| 569 | + ICollection inner = _captures.Values; // c[1] | ||
| 570 | 570 | inner.CopyTo(captures, 0); | |
| 571 | 571 | return captures; | |
| 572 | 572 | } | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -340,7 +340,7 @@ Tensors run_internal_graph(Tensors inputs, bool training = false, Tensors mask = | |||
| 340 | 340 | tf.Logger.Debug($"{depth}: {node.Layer}: {node.Layer.Name}"); | |
| 341 | 341 | var outputs = node.Layer.Apply(layer_inputs, is_training: training); | |
| 342 | 342 | foreach (var output in outputs.Where(x => x != null)) | |
| 343 | - tf.Logger.Debug($"{output.TensorShape}"); | ||
| 343 | + tf.Logger.Debug($"{depth}: {node.Layer}: {node.Layer.Name} {output.TensorShape}"); | ||
| 344 | 344 | // Update tensor_dict for next input | |
| 345 | 345 | foreach (var (x_id, y) in zip(node.FlatOutputIds, outputs)) | |
| 346 | 346 | tensor_dict[x_id] = new Queue<Tensor>(Enumerable.Range(0, tensor_usage_count[x_id]).Select(x => y)); | |
| Back | FazBrowse Home | New Git URL |
0 commit comments