| FazBrowse GitHub Viewer | Trending | | Home |
| Tools: [Download Repo ZIP] [Original HTTPS Page] |
1 parent c93c219 commit 67a70bf
6 files changed
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -2,7 +2,9 @@ | |||
| 2 | 2 | using System.Collections.Generic; | |
| 3 | 3 | using System.Linq; | |
| 4 | 4 | using System.Text; | |
| 5 | + using Tensorflow.Eager; | ||
| 5 | 6 | using Tensorflow.Graphs; | |
| 7 | + using Tensorflow.NumPy; | ||
| 6 | 8 | using static Tensorflow.Binding; | |
| 7 | 9 | using static Tensorflow.tensorflow; | |
| 8 | 10 | ||
@@ -148,7 +150,7 @@ public void Record(Tensors flat_outputs, Tensors inference_args) | |||
| 148 | 150 | src_graph: _func_graph); | |
| 149 | 151 | ||
| 150 | 152 | var captures_from_forward = backwards_graph.external_captures | |
| 151 | - .Where(x => x.IsCreatedInGraphMode && x.graph == _func_graph) | ||
| 153 | + .Where(x => x is not EagerTensor && x is not NDArray && x.graph == _func_graph) | ||
| 152 | 154 | .ToArray(); | |
| 153 | 155 | foreach(var capture in captures_from_forward) | |
| 154 | 156 | { | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -32,7 +32,6 @@ public partial class Tensor | |||
| 32 | 32 | ||
| 33 | 33 | public Tensor() | |
| 34 | 34 | { | |
| 35 | - _isCreatedInGraphMode = !tf.executing_eagerly(); | ||
| 36 | 35 | } | |
| 37 | 36 | ||
| 38 | 37 | /// <summary> | |
@@ -44,8 +43,6 @@ public unsafe Tensor(SafeTensorHandle handle, bool clone = false) | |||
| 44 | 43 | _handle = handle; | |
| 45 | 44 | if (clone && handle != null) | |
| 46 | 45 | _handle = TF_NewTensor(shape, dtype, data: TensorDataPointer.ToPointer()); | |
| 47 | - | ||
| 48 | - _isCreatedInGraphMode = !tf.executing_eagerly(); | ||
| 49 | 46 | } | |
| 50 | 47 | ||
| 51 | 48 | /// <summary> | |
@@ -59,13 +56,11 @@ public unsafe Tensor(SafeTensorHandle handle, bool clone = false) | |||
| 59 | 56 | public unsafe Tensor(IntPtr data_ptr, Shape shape, TF_DataType dtype) | |
| 60 | 57 | { | |
| 61 | 58 | _handle = TF_NewTensor(shape, dtype, data: data_ptr.ToPointer()); | |
| 62 | - _isCreatedInGraphMode = !tf.executing_eagerly(); | ||
| 63 | 59 | } | |
| 64 | 60 | ||
| 65 | 61 | public unsafe Tensor(NDArray nd) | |
| 66 | 62 | { | |
| 67 | 63 | _handle = TF_NewTensor(nd.shape, nd.dtype, nd.data.ToPointer()); | |
| 68 | - _isCreatedInGraphMode = !tf.executing_eagerly(); | ||
| 69 | 64 | } | |
| 70 | 65 | ||
| 71 | 66 | #region scala | |
@@ -107,13 +102,11 @@ public Tensor(Operation op, int value_index, TF_DataType dtype) | |||
| 107 | 102 | _value_index = value_index; | |
| 108 | 103 | _override_dtype = dtype; | |
| 109 | 104 | _id = ops.uid(); | |
| 110 | - _isCreatedInGraphMode = !tf.executing_eagerly(); | ||
| 111 | 105 | } | |
| 112 | 106 | ||
| 113 | 107 | protected unsafe void InitTensor(Shape shape, TF_DataType dtype) | |
| 114 | 108 | { | |
| 115 | 109 | _handle = TF_NewTensor(shape, dtype, null); | |
| 116 | - _isCreatedInGraphMode = !tf.executing_eagerly(); | ||
| 117 | 110 | } | |
| 118 | 111 | ||
| 119 | 112 | protected unsafe void InitTensor(Shape shape, byte[] bytes, TF_DataType dtype) | |
@@ -122,13 +115,10 @@ protected unsafe void InitTensor(Shape shape, byte[] bytes, TF_DataType dtype) | |||
| 122 | 115 | _handle = StringTensor(new byte[][] { bytes }, Shape.Scalar); | |
| 123 | 116 | else | |
| 124 | 117 | _handle = TF_NewTensor(bytes, shape, dtype); | |
| 125 | - _isCreatedInGraphMode = !tf.executing_eagerly(); | ||
| 126 | 118 | } | |
| 127 | 119 | ||
| 128 | 120 | protected unsafe void InitTensor(Array array, Shape? shape = null) | |
| 129 | 121 | { | |
| 130 | - _isCreatedInGraphMode = !tf.executing_eagerly(); | ||
| 131 | - | ||
| 132 | 122 | shape = shape ?? array.GetShape(); | |
| 133 | 123 | var dtype = array.GetDataType(); | |
| 134 | 124 | ||
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -94,10 +94,6 @@ public partial class Tensor : DisposableObject, | |||
| 94 | 94 | /// </summary> | |
| 95 | 95 | public SafeEagerTensorHandle EagerTensorHandle => _eagerTensorHandle; | |
| 96 | 96 | ||
| 97 | - protected bool _isCreatedInGraphMode; | ||
| 98 | - | ||
| 99 | - public bool IsCreatedInGraphMode => _isCreatedInGraphMode; | ||
| 100 | - | ||
| 101 | 97 | /// <summary> | |
| 102 | 98 | /// Returns the shape of a tensor. | |
| 103 | 99 | /// </summary> | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -21,7 +21,6 @@ public class Tensors : IEnumerable<Tensor>, IDisposable | |||
| 21 | 21 | public Shape shape => items.First().shape; | |
| 22 | 22 | public int rank => items.First().rank; | |
| 23 | 23 | public Graph graph => items.First().graph; | |
| 24 | - public bool IsCreatedInGraphMode => items.First().IsCreatedInGraphMode; | ||
| 25 | 24 | public bool IsList { get; set; } | |
| 26 | 25 | public int Length => items.Count(); | |
| 27 | 26 | ||
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -68,7 +68,7 @@ public void __init__(bool trainable = true, | |||
| 68 | 68 | // when this object is garbage collected the deleter will be too. This | |
| 69 | 69 | // means ResourceVariables can be part of reference cycles without those | |
| 70 | 70 | // cycles being uncollectable. | |
| 71 | - if (!handle.IsCreatedInGraphMode) | ||
| 71 | + if (handle is EagerTensor) | ||
| 72 | 72 | { | |
| 73 | 73 | _handle = handle.EagerTensorHandle.DangerousGetHandle(); | |
| 74 | 74 | eager_resource_deleter = new EagerResourceDeleter(handle, handle.Device); | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -18,9 +18,11 @@ limitations under the License. | |||
| 18 | 18 | using System.Collections.Generic; | |
| 19 | 19 | using System.Linq; | |
| 20 | 20 | using System.Threading; | |
| 21 | + using Tensorflow.Eager; | ||
| 21 | 22 | using Tensorflow.Keras.ArgsDefinition; | |
| 22 | 23 | using Tensorflow.Keras.Saving; | |
| 23 | 24 | using Tensorflow.Keras.Utils; | |
| 25 | + using Tensorflow.NumPy; | ||
| 24 | 26 | using Tensorflow.Train; | |
| 25 | 27 | using static Tensorflow.Binding; | |
| 26 | 28 | ||
@@ -118,7 +120,7 @@ public Layer(LayerArgs args) | |||
| 118 | 120 | bool _in_functional_construction_mode(Tensors inputs) | |
| 119 | 121 | { | |
| 120 | 122 | return tf.Context.executing_eagerly() | |
| 121 | - && inputs.Count(x => x.IsCreatedInGraphMode) == inputs.Count(); | ||
| 123 | + && inputs.Count(x => x is not EagerTensor && x is not NDArray) == inputs.Count(); | ||
| 122 | 124 | } | |
| 123 | 125 | ||
| 124 | 126 | public void SetConnectivityMetadata(Tensors inputs, Tensors outputs) | |
@@ -180,7 +182,7 @@ protected void MaybeBuild(Tensors inputs) | |||
| 180 | 182 | tf.init_scope(); | |
| 181 | 183 | ||
| 182 | 184 | bool need_restore_mode = false; | |
| 183 | - if (!inputs.IsCreatedInGraphMode || tf.Context.is_build_function()) | ||
| 185 | + if (inputs.Any(x => x is EagerTensor) || tf.Context.is_build_function()) | ||
| 184 | 186 | { | |
| 185 | 187 | need_restore_mode = true; | |
| 186 | 188 | tf.Context.eager_mode(isFunc: tf.Context.is_build_function()); | |
| Back | FazBrowse Home | New Git URL |
0 commit comments