| FazBrowse GitHub Viewer | Trending | | Home |
| Tools: [Download Repo ZIP] [Original HTTPS Page] |
1 parent 05443ea commit 1460ec8
6 files changed
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -13,6 +13,7 @@ namespace Tensorflow.Functions | |||
| 13 | 13 | public class ConcreteFunction | |
| 14 | 14 | { | |
| 15 | 15 | FuncGraph func_graph; | |
| 16 | + ForwardBackwardCall forward_backward; | ||
| 16 | 17 | public Tensor[] Inputs => func_graph.Inputs; | |
| 17 | 18 | public Tensor[] CapturedInputs => func_graph.external_captures; | |
| 18 | 19 | ||
@@ -151,7 +152,8 @@ public Tensors CallFlat(Tensor[] args, Tensor[] captured_inputs) | |||
| 151 | 152 | return tf.Runner.Execute(tf.Context, func_graph.FuncName, func_graph.Outputs.Length, args, attrs); | |
| 152 | 153 | } | |
| 153 | 154 | ||
| 154 | - var forward_backward = SelectForwardAndBackwardFunctions(args, possible_gradient_type, executing_eagerly); | ||
| 155 | + if (forward_backward == null) | ||
| 156 | + forward_backward = SelectForwardAndBackwardFunctions(args, possible_gradient_type, executing_eagerly); | ||
| 155 | 157 | var (forward_function, args_with_tangents) = forward_backward.Forward(); | |
| 156 | 158 | Tensors flat_outputs = null; | |
| 157 | 159 | if (executing_eagerly) | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -13,6 +13,7 @@ public class ForwardBackwardCall | |||
| 13 | 13 | Tensors _inference_args; | |
| 14 | 14 | Tensors _input_tangents; | |
| 15 | 15 | bool _tape_watching; | |
| 16 | + EagerDefinedFunction forward_function; | ||
| 16 | 17 | ||
| 17 | 18 | public ForwardBackwardCall(TapeGradientFunctions functions, | |
| 18 | 19 | Tensors inference_args, | |
@@ -22,10 +23,11 @@ public ForwardBackwardCall(TapeGradientFunctions functions, | |||
| 22 | 23 | _inference_args = inference_args; | |
| 23 | 24 | _tape_watching = tape_watching; | |
| 24 | 25 | } | |
| 25 | - | ||
| 26 | + | ||
| 26 | 27 | public (EagerDefinedFunction, Tensors) Forward() | |
| 27 | 28 | { | |
| 28 | - var forward_function = _functions.Forward(_inference_args); | ||
| 29 | + if (forward_function == null) | ||
| 30 | + forward_function = _functions.Forward(_inference_args); | ||
| 29 | 31 | return (forward_function, _inference_args); | |
| 30 | 32 | } | |
| 31 | 33 | ||
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -25,6 +25,7 @@ public abstract class TapeGradientFunctions | |||
| 25 | 25 | protected List<int> _forwardprop_output_indices; | |
| 26 | 26 | protected int _num_forwardprop_outputs; | |
| 27 | 27 | protected ConcreteFunction _backward; | |
| 28 | + BackwardFunction _backward_function_wrapper; | ||
| 28 | 29 | ||
| 29 | 30 | public TapeGradientFunctions(FuncGraph func_graph, | |
| 30 | 31 | bool need_gradients_for_jvps) | |
@@ -58,60 +59,66 @@ public void Record(Tensors flat_outputs, Tensors inference_args) | |||
| 58 | 59 | /// <returns></returns> | |
| 59 | 60 | (BackwardFunction, Tensors) _wrap_backward_function(FuncGraph forward_graph, ConcreteFunction backward, Tensors outputs) | |
| 60 | 61 | { | |
| 61 | - var capture_mapping = new Dictionary<long, Tensor>(); | ||
| 62 | - foreach(var (i, output) in enumerate(outputs)) | ||
| 63 | - capture_mapping[forward_graph.Outputs[i].Id] = output; | ||
| 64 | - | ||
| 65 | - var remapped_captures = new Tensors(); | ||
| 66 | - foreach(var capture in backward.CapturedInputs) | ||
| 67 | - { | ||
| 68 | - if (capture_mapping.ContainsKey(capture.Id)) | ||
| 69 | - remapped_captures.Add(capture_mapping[capture.Id]); | ||
| 70 | - } | ||
| 71 | - | ||
| 72 | 62 | var backward_function_inputs = backward.Inputs.Length - backward.CapturedInputs.Length; | |
| 73 | 63 | var recorded_outputs = new Tensors(); | |
| 74 | - var relevant_outputs = outputs; | ||
| 75 | 64 | var trainable_recorded_outputs = 0; | |
| 76 | - var skip_positions = new List<int>(); | ||
| 77 | - foreach (var (output_index, output) in enumerate(relevant_outputs)) | ||
| 65 | + foreach (var (output_index, output) in enumerate(outputs)) | ||
| 78 | 66 | { | |
| 79 | 67 | if (trainable_recorded_outputs < backward_function_inputs) | |
| 80 | 68 | recorded_outputs.Add(output); | |
| 81 | 69 | if (gradients_util.IsTrainable(output)) | |
| 82 | 70 | trainable_recorded_outputs += 1; | |
| 83 | - else | ||
| 84 | - skip_positions.Add(output_index); | ||
| 85 | 71 | } | |
| 86 | 72 | ||
| 87 | - BackwardFunction _backward_function_wrapper = (args, unneeded_gradients) => | ||
| 73 | + if(_backward_function_wrapper == null) | ||
| 88 | 74 | { | |
| 89 | - var processed_args = new Tensors(); | ||
| 90 | - var input_index = 0; | ||
| 91 | - foreach (var (output_index, arg) in enumerate(args)) | ||
| 75 | + var capture_mapping = new Dictionary<long, Tensor>(); | ||
| 76 | + foreach (var (i, output) in enumerate(outputs)) | ||
| 77 | + capture_mapping[forward_graph.Outputs[i].Id] = output; | ||
| 78 | + | ||
| 79 | + var remapped_captures = new Tensors(); | ||
| 80 | + foreach (var capture in backward.CapturedInputs) | ||
| 92 | 81 | { | |
| 93 | - if (skip_positions.Contains(output_index)) | ||
| 94 | - continue; | ||
| 95 | - if (arg == null) | ||
| 96 | - throw new NotImplementedException(""); | ||
| 97 | - processed_args.Add(arg); | ||
| 98 | - input_index += 1; | ||
| 99 | - if (input_index >= backward_function_inputs) | ||
| 100 | - break; | ||
| 82 | + if (capture_mapping.ContainsKey(capture.Id)) | ||
| 83 | + remapped_captures.Add(capture_mapping[capture.Id]); | ||
| 101 | 84 | } | |
| 102 | 85 | ||
| 103 | - tf.Logger.Debug($"Invoke backward function: {backward.Name}"); | ||
| 104 | - var gradients = backward.CallFlat(processed_args, remapped_captures); | ||
| 105 | - | ||
| 106 | - foreach (var unneeded_gradient_index in unneeded_gradients) | ||
| 86 | + var skip_positions = new List<int>(); | ||
| 87 | + foreach (var (output_index, output) in enumerate(outputs)) | ||
| 107 | 88 | { | |
| 108 | - var index = Convert.ToInt32(unneeded_gradient_index); | ||
| 109 | - if (gradients.Length <= index) | ||
| 110 | - gradients.Insert(index, null); | ||
| 89 | + if (!gradients_util.IsTrainable(output)) | ||
| 90 | + skip_positions.Add(output_index); | ||
| 111 | 91 | } | |
| 112 | 92 | ||
| 113 | - return gradients; | ||
| 114 | - }; | ||
| 93 | + _backward_function_wrapper = (args, unneeded_gradients) => | ||
| 94 | + { | ||
| 95 | + var processed_args = new Tensors(); | ||
| 96 | + var input_index = 0; | ||
| 97 | + foreach (var (output_index, arg) in enumerate(args)) | ||
| 98 | + { | ||
| 99 | + if (skip_positions.Contains(output_index)) | ||
| 100 | + continue; | ||
| 101 | + if (arg == null) | ||
| 102 | + throw new NotImplementedException(""); | ||
| 103 | + processed_args.Add(arg); | ||
| 104 | + input_index += 1; | ||
| 105 | + if (input_index >= backward_function_inputs) | ||
| 106 | + break; | ||
| 107 | + } | ||
| 108 | + | ||
| 109 | + tf.Logger.Debug($"Invoke backward function: {backward.Name}"); | ||
| 110 | + var gradients = backward.CallFlat(processed_args, remapped_captures); | ||
| 111 | + | ||
| 112 | + foreach (var unneeded_gradient_index in unneeded_gradients) | ||
| 113 | + { | ||
| 114 | + var index = Convert.ToInt32(unneeded_gradient_index); | ||
| 115 | + if (gradients.Length <= index) | ||
| 116 | + gradients.Insert(index, null); | ||
| 117 | + } | ||
| 118 | + | ||
| 119 | + return gradients; | ||
| 120 | + }; | ||
| 121 | + } | ||
| 115 | 122 | ||
| 116 | 123 | return (_backward_function_wrapper, recorded_outputs); | |
| 117 | 124 | } | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -376,6 +376,10 @@ public static int GraphUniqueId() | |||
| 376 | 376 | public static int uid_function() | |
| 377 | 377 | => Interlocked.Increment(ref uid_number_for_function); | |
| 378 | 378 | ||
| 379 | + static int uid_number_for_layer = 0; | ||
| 380 | + public static int uid_layer() | ||
| 381 | + => Interlocked.Increment(ref uid_number_for_layer); | ||
| 382 | + | ||
| 379 | 383 | public static void reset_uid() | |
| 380 | 384 | { | |
| 381 | 385 | uid_number = -1; | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -66,6 +66,8 @@ public abstract partial class Layer : AutoTrackable, ILayer | |||
| 66 | 66 | protected List<IVariableV1> non_trainable_weights; | |
| 67 | 67 | public List<IVariableV1> non_trainable_variables => non_trainable_weights; | |
| 68 | 68 | ||
| 69 | + protected int id; | ||
| 70 | + public int Id => id; | ||
| 69 | 71 | protected string name; | |
| 70 | 72 | protected string base_name; | |
| 71 | 73 | public string Name => name; | |
@@ -96,6 +98,7 @@ public Layer(LayerArgs args) | |||
| 96 | 98 | built = false; | |
| 97 | 99 | SupportsMasking = false; | |
| 98 | 100 | ||
| 101 | + id = ops.uid_layer(); | ||
| 99 | 102 | _init_set_name(args.Name); | |
| 100 | 103 | trainable_weights = new List<IVariableV1>(); | |
| 101 | 104 | non_trainable_weights = new List<IVariableV1>(); | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -8,6 +8,7 @@ | |||
| 8 | 8 | using Tensorflow.Keras.ArgsDefinition; | |
| 9 | 9 | using Tensorflow.Keras.Engine; | |
| 10 | 10 | using static Tensorflow.Binding; | |
| 11 | + using Tensorflow.Functions; | ||
| 11 | 12 | ||
| 12 | 13 | namespace Tensorflow.Keras.Layers | |
| 13 | 14 | { | |
@@ -35,10 +36,39 @@ public TensorFlowOpLayer(TensorFlowOpLayerArgs args) | |||
| 35 | 36 | protected override Tensors Call(Tensors inputs, Tensor state = null, bool? training = null) | |
| 36 | 37 | { | |
| 37 | 38 | if (tf.Context.executing_eagerly()) | |
| 38 | - return _defun_call(inputs); | ||
| 39 | + return DeFunCall(inputs); | ||
| 39 | 40 | return MakOp(inputs); | |
| 40 | 41 | } | |
| 41 | 42 | ||
| 43 | + ConcreteFunction function; | ||
| 44 | + Tensors DeFunCall(Tensors inputs) | ||
| 45 | + { | ||
| 46 | + if(function == null) | ||
| 47 | + { | ||
| 48 | + function = new ConcreteFunction(name); | ||
| 49 | + function.Enter(); | ||
| 50 | + | ||
| 51 | + int i = 0; | ||
| 52 | + var graph_inputs = inputs.Select(x => tf.placeholder(x.dtype, shape: x.shape, name: $"defun_inputs_{i++}")).ToArray(); | ||
| 53 | + var graph_outputs = MakOp(graph_inputs); | ||
| 54 | + graph_outputs = mark_as_return(graph_outputs); | ||
| 55 | + | ||
| 56 | + function.ToGraph(graph_inputs, graph_outputs); | ||
| 57 | + function.Exit(); | ||
| 58 | + } | ||
| 59 | + | ||
| 60 | + var outputs = function.FilteredCall(inputs); | ||
| 61 | + return outputs; | ||
| 62 | + } | ||
| 63 | + | ||
| 64 | + Tensors mark_as_return(Tensors tensors) | ||
| 65 | + { | ||
| 66 | + var result = new Tensors(); | ||
| 67 | + foreach (var tensor in tensors) | ||
| 68 | + result.Add(array_ops.identity(tensor)); | ||
| 69 | + return result; | ||
| 70 | + } | ||
| 71 | + | ||
| 42 | 72 | [AutoGraph] | |
| 43 | 73 | Tensors _defun_call(Tensors inputs) | |
| 44 | 74 | => MakOp(inputs); | |
| Back | FazBrowse Home | New Git URL |
0 commit comments