| FazBrowse GitHub Viewer | Trending | | Home |
| Tools: [Download Repo ZIP] [Original HTTPS Page] |
1 parent ceccf40 commit 1ae9bbc
24 files changed
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -78,15 +78,15 @@ public Tensor check_numerics(Tensor tensor, string message, string name = null) | |||
| 78 | 78 | /// <param name="axis"></param> | |
| 79 | 79 | /// <param name="name"></param> | |
| 80 | 80 | /// <returns>A `Tensor` resulting from concatenation of the input tensors.</returns> | |
| 81 | - public Tensor concat(IList<Tensor> values, int axis, string name = "concat") | ||
| 81 | + public Tensor concat(IEnumerable<Tensor> values, int axis, string name = "concat") | ||
| 82 | 82 | { | |
| 83 | - if (values.Count == 1) | ||
| 83 | + if (values.Count() == 1) | ||
| 84 | 84 | { | |
| 85 | 85 | return tf_with(ops.name_scope(name), scope => | |
| 86 | 86 | { | |
| 87 | 87 | var tensor = ops.convert_to_tensor(axis, name: "concat_dim", dtype: dtypes.int32); | |
| 88 | 88 | Debug.Assert(tensor.TensorShape.ndim == 0); | |
| 89 | - return identity(values[0], name: scope); | ||
| 89 | + return identity(values.First(), name: scope); | ||
| 90 | 90 | }); | |
| 91 | 91 | } | |
| 92 | 92 | ||
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -19,15 +19,18 @@ namespace Tensorflow | |||
| 19 | 19 | public partial class tensorflow | |
| 20 | 20 | { | |
| 21 | 21 | public Tensor reshape(Tensor tensor, | |
| 22 | - TensorShape shape, | ||
| 23 | - string name = null) => gen_array_ops.reshape(tensor, shape, name); | ||
| 22 | + TensorShape shape, | ||
| 23 | + string name = null) | ||
| 24 | + => gen_array_ops.reshape(tensor, shape, name); | ||
| 24 | 25 | ||
| 25 | 26 | public Tensor reshape(Tensor tensor, | |
| 26 | - Tensor[] shape, | ||
| 27 | - string name = null) => gen_array_ops.reshape(tensor, shape, name); | ||
| 27 | + Tensor shape, | ||
| 28 | + string name = null) | ||
| 29 | + => gen_array_ops.reshape(tensor, shape, name); | ||
| 28 | 30 | ||
| 29 | 31 | public Tensor reshape(Tensor tensor, | |
| 30 | - Tensor shape, | ||
| 31 | - string name = null) => gen_array_ops.reshape(tensor, shape, name); | ||
| 32 | + object[] shape, | ||
| 33 | + string name = null) | ||
| 34 | + => gen_array_ops.reshape(tensor, shape, name); | ||
| 32 | 35 | } | |
| 33 | 36 | } | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -13,13 +13,22 @@ You may obtain a copy of the License at | |||
| 13 | 13 | See the License for the specific language governing permissions and | |
| 14 | 14 | limitations under the License. | |
| 15 | 15 | ******************************************************************************/ | |
| 16 | + using static Tensorflow.Binding; | ||
| 16 | 17 | ||
| 17 | 18 | namespace Tensorflow | |
| 18 | 19 | { | |
| 19 | 20 | public partial class tensorflow | |
| 20 | 21 | { | |
| 21 | - public Tensor tile<T>(Tensor input, | ||
| 22 | - T multiples, | ||
| 23 | - string name = null) => gen_array_ops.tile(input, multiples, name); | ||
| 22 | + public Tensor tile(Tensor input, Tensor multiples, string name = null) | ||
| 23 | + => gen_array_ops.tile(input, multiples, name); | ||
| 24 | + | ||
| 25 | + public Tensor tile(Tensor input, object[] multiples, string name = null) | ||
| 26 | + => gen_array_ops.tile(input, multiples, name); | ||
| 27 | + | ||
| 28 | + public Tensor tile(Tensor input, TensorShape multiples, string name = null) | ||
| 29 | + { | ||
| 30 | + var multiples_tensor = constant_op.constant(multiples); | ||
| 31 | + return gen_array_ops.tile(input, multiples_tensor, name); | ||
| 32 | + } | ||
| 24 | 33 | } | |
| 25 | 34 | } | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -28,16 +28,27 @@ namespace Tensorflow.Contexts | |||
| 28 | 28 | /// </summary> | |
| 29 | 29 | public sealed partial class Context | |
| 30 | 30 | { | |
| 31 | - // [DebuggerStepThrough] | ||
| 32 | - public T RunInAutoMode<T>(Func<T> graphAction, Func<T> eagerAction, params Tensor[] tensors) | ||
| 31 | + public T RunInAutoMode<T>(Func<T> graphAction, Func<T> eagerAction, params object[] args) | ||
| 33 | 32 | { | |
| 34 | - var shouldRunInEager = executing_eagerly() | ||
| 35 | - && tensors.Count(x => x.IsEagerTensor) == tensors.Length; | ||
| 36 | - | ||
| 37 | - if (shouldRunInEager) | ||
| 38 | - return eagerAction(); | ||
| 39 | - else | ||
| 33 | + if (tf.Context.has_graph_arg(args)) | ||
| 34 | + { | ||
| 40 | 35 | return graphAction(); | |
| 36 | + } | ||
| 37 | + else | ||
| 38 | + { | ||
| 39 | + try | ||
| 40 | + { | ||
| 41 | + return eagerAction(); | ||
| 42 | + } | ||
| 43 | + catch (InvalidArgumentError ex) | ||
| 44 | + { | ||
| 45 | + throw ex; | ||
| 46 | + } | ||
| 47 | + catch (Exception ex) | ||
| 48 | + { | ||
| 49 | + return graphAction(); | ||
| 50 | + } | ||
| 51 | + } | ||
| 41 | 52 | } | |
| 42 | 53 | ||
| 43 | 54 | // [DebuggerStepThrough] | |
@@ -46,12 +57,7 @@ public Tensors RunInAutoMode2(Func<Tensors> graphAction, | |||
| 46 | 57 | Action<Operation> recordGradient, | |
| 47 | 58 | Tensors tensors) | |
| 48 | 59 | { | |
| 49 | - var shouldRunInEager = executing_eagerly() | ||
| 50 | - && tensors.Count(x => x.IsEagerTensor) == tensors.Length; | ||
| 51 | - | ||
| 52 | - if (shouldRunInEager) | ||
| 53 | - return eagerAction(); | ||
| 54 | - else | ||
| 60 | + if (tf.Context.has_graph_arg(tensors)) | ||
| 55 | 61 | { | |
| 56 | 62 | if (executing_eagerly()) | |
| 57 | 63 | { | |
@@ -68,6 +74,10 @@ public Tensors RunInAutoMode2(Func<Tensors> graphAction, | |||
| 68 | 74 | return result; | |
| 69 | 75 | } | |
| 70 | 76 | } | |
| 77 | + else | ||
| 78 | + { | ||
| 79 | + return eagerAction(); | ||
| 80 | + } | ||
| 71 | 81 | } | |
| 72 | 82 | } | |
| 73 | 83 | } | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -20,6 +20,7 @@ limitations under the License. | |||
| 20 | 20 | using Tensorflow.Eager; | |
| 21 | 21 | using static Tensorflow.Binding; | |
| 22 | 22 | using Google.Protobuf; | |
| 23 | + using Tensorflow.Util; | ||
| 23 | 24 | ||
| 24 | 25 | namespace Tensorflow.Contexts | |
| 25 | 26 | { | |
@@ -103,6 +104,29 @@ public void graph_mode(bool isFunc = false) | |||
| 103 | 104 | public void eager_mode(bool isFunc = false) | |
| 104 | 105 | => context_switches.Push(true, isFunc); | |
| 105 | 106 | ||
| 107 | + public bool switched_to_graph(params object[] args) | ||
| 108 | + { | ||
| 109 | + var switching_to_graph = has_graph_arg(args) && tf.Context.executing_eagerly(); | ||
| 110 | + if (switching_to_graph) | ||
| 111 | + tf.Context.graph_mode(tf.Context.is_build_function()); | ||
| 112 | + return switching_to_graph; | ||
| 113 | + } | ||
| 114 | + | ||
| 115 | + public bool has_graph_arg(params object[] args) | ||
| 116 | + { | ||
| 117 | + var flatten_args = nest.flatten<object>(args); | ||
| 118 | + bool has_graph_arg = false; | ||
| 119 | + foreach (var el in flatten_args) | ||
| 120 | + { | ||
| 121 | + if (el is Tensor tensor && !tensor.IsEagerTensor) | ||
| 122 | + { | ||
| 123 | + has_graph_arg = true; | ||
| 124 | + break; | ||
| 125 | + } | ||
| 126 | + } | ||
| 127 | + return has_graph_arg; | ||
| 128 | + } | ||
| 129 | + | ||
| 106 | 130 | public void restore_mode() | |
| 107 | 131 | { | |
| 108 | 132 | context_switches.Pop(); | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -38,9 +38,9 @@ public bool RecordGradient(string op_name, | |||
| 38 | 38 | } | |
| 39 | 39 | }*/ | |
| 40 | 40 | } | |
| 41 | - | ||
| 42 | - tf.Logger.Debug($"RecordGradient: should_record={should_record}, op_name={op_name}"); | ||
| 41 | + | ||
| 43 | 42 | if (!should_record) return should_record; | |
| 43 | + tf.Logger.Debug($"RecordGradient: op_name={op_name}"); | ||
| 44 | 44 | ||
| 45 | 45 | Tensor[] op_outputs; | |
| 46 | 46 | #pragma warning disable CS0219 // Variable is assigned but its value is never used | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -50,7 +50,7 @@ public Tensor[] TFE_FastPathExecute(Context ctx, | |||
| 50 | 50 | ||
| 51 | 51 | var op_def = tf.get_default_graph().GetOpDef(opName); | |
| 52 | 52 | ||
| 53 | - var flattened_attrs = new List<object>(op_def.InputArg.Count); | ||
| 53 | + var flattened_attrs = new List<object>(op_def.Attr.Count * 2); | ||
| 54 | 54 | var flattened_inputs = new List<Tensor>(op_def.InputArg.Count); | |
| 55 | 55 | ||
| 56 | 56 | // Set non-inferred attrs, including setting defaults if the attr is passed in | |
@@ -221,23 +221,9 @@ bool AddInputToOp(object inputs, | |||
| 221 | 221 | SafeTensorHandleHandle input_handle; | |
| 222 | 222 | ||
| 223 | 223 | // ConvertToTensor(); | |
| 224 | - switch (inputs) | ||
| 225 | - { | ||
| 226 | - case EagerTensor input: | ||
| 227 | - input_handle = input.EagerTensorHandle; | ||
| 228 | - flattened_inputs.Add(input); | ||
| 229 | - break; | ||
| 230 | - case ResourceVariable variable: | ||
| 231 | - var var_tensor = variable.AsTensor(); | ||
| 232 | - input_handle = var_tensor.EagerTensorHandle; | ||
| 233 | - flattened_inputs.Add(var_tensor); | ||
| 234 | - break; | ||
| 235 | - default: | ||
| 236 | - var tensor = tf.convert_to_tensor(inputs); | ||
| 237 | - input_handle = tensor.EagerTensorHandle; | ||
| 238 | - flattened_inputs.Add(tensor); | ||
| 239 | - break; | ||
| 240 | - } | ||
| 224 | + var tensor = tf.convert_to_tensor(inputs); | ||
| 225 | + input_handle = tensor.EagerTensorHandle; | ||
| 226 | + flattened_inputs.Add(tensor); | ||
| 241 | 227 | ||
| 242 | 228 | if (add_type_attr && !string.IsNullOrEmpty(input_arg.TypeAttr)) | |
| 243 | 229 | { | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -2,6 +2,7 @@ | |||
| 2 | 2 | using System; | |
| 3 | 3 | using System.Linq; | |
| 4 | 4 | using System.Text; | |
| 5 | + using static Tensorflow.Binding; | ||
| 5 | 6 | ||
| 6 | 7 | namespace Tensorflow.Framework | |
| 7 | 8 | { | |
@@ -65,5 +66,17 @@ public static int dimension_value(Dimension dimension) | |||
| 65 | 66 | ||
| 66 | 67 | public static TensorShape as_shape(this Shape shape) | |
| 67 | 68 | => new TensorShape(shape.Dimensions); | |
| 69 | + | ||
| 70 | + public static TensorShape most_specific_compatible_shape(this TensorShape self, TensorShape other) | ||
| 71 | + { | ||
| 72 | + var dims = range(self.rank).Select(x => -1).ToArray(); | ||
| 73 | + foreach(var (i, (d1, d2)) in enumerate(zip(self.dims, other.dims))) | ||
| 74 | + { | ||
| 75 | + if (d1 == d2) | ||
| 76 | + dims[i] = d1; | ||
| 77 | + } | ||
| 78 | + | ||
| 79 | + return new TensorShape(dims); | ||
| 80 | + } | ||
| 68 | 81 | } | |
| 69 | 82 | } | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -134,7 +134,7 @@ public Tensors FilteredCall(Tensors inputs) | |||
| 134 | 134 | /// <param name="args"></param> | |
| 135 | 135 | /// <param name="captured_inputs"></param> | |
| 136 | 136 | /// <returns></returns> | |
| 137 | - public Tensor[] CallFlat(Tensor[] args, Tensor[] captured_inputs) | ||
| 137 | + public Tensors CallFlat(Tensor[] args, Tensor[] captured_inputs) | ||
| 138 | 138 | { | |
| 139 | 139 | var executing_eagerly = tf.Context.executing_eagerly(); | |
| 140 | 140 | var default_graph = ops.get_default_graph(); | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -99,8 +99,18 @@ public void Record(Tensors flat_outputs, Tensors inference_args) | |||
| 99 | 99 | if (input_index >= backward_function_inputs) | |
| 100 | 100 | break; | |
| 101 | 101 | } | |
| 102 | + | ||
| 102 | 103 | tf.Logger.Debug($"Invoke backward function: {backward.Name}"); | |
| 103 | - return backward.CallFlat(processed_args, remapped_captures); | ||
| 104 | + var gradients = backward.CallFlat(processed_args, remapped_captures); | ||
| 105 | + | ||
| 106 | + foreach (var unneeded_gradient_index in unneeded_gradients) | ||
| 107 | + { | ||
| 108 | + var index = Convert.ToInt32(unneeded_gradient_index); | ||
| 109 | + if (gradients.Length <= index) | ||
| 110 | + gradients.Insert(index, null); | ||
| 111 | + } | ||
| 112 | + | ||
| 113 | + return gradients; | ||
| 104 | 114 | }; | |
| 105 | 115 | ||
| 106 | 116 | return (_backward_function_wrapper, recorded_outputs); | |
| Back | FazBrowse Home | New Git URL |
0 commit comments