| FazBrowse GitHub Viewer | Trending | | Home |
| Tools: [Download Repo ZIP] [Original HTTPS Page] |
1 parent 71488b4 commit 39883ae
9 files changed
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -10,7 +10,7 @@ public partial class EagerTensor | |||
| 10 | 10 | public EagerTensor(SafeTensorHandleHandle handle) | |
| 11 | 11 | { | |
| 12 | 12 | _id = ops.uid(); | |
| 13 | - EagerTensorHandle = handle; | ||
| 13 | + _eagerTensorHandle = handle; | ||
| 14 | 14 | Resolve(); | |
| 15 | 15 | } | |
| 16 | 16 | ||
@@ -59,20 +59,14 @@ public EagerTensor(byte[] bytes, TF_DataType dtype) : base(bytes, dtype) | |||
| 59 | 59 | void NewEagerTensorHandle(IntPtr h) | |
| 60 | 60 | { | |
| 61 | 61 | _id = ops.uid(); | |
| 62 | - EagerTensorHandle = c_api.TFE_NewTensorHandle(h, tf.Status.Handle); | ||
| 62 | + _eagerTensorHandle = c_api.TFE_NewTensorHandle(h, tf.Status.Handle); | ||
| 63 | 63 | tf.Status.Check(true); | |
| 64 | - #if TRACK_TENSOR_LIFE | ||
| 65 | - print($"New EagerTensorHandle {EagerTensorHandle} {Id} From 0x{h.ToString("x16")}"); | ||
| 66 | - #endif | ||
| 67 | 64 | } | |
| 68 | 65 | ||
| 69 | 66 | private void Resolve() | |
| 70 | 67 | { | |
| 71 | - _handle = c_api.TFE_TensorHandleResolve(EagerTensorHandle, tf.Status.Handle); | ||
| 68 | + _handle = c_api.TFE_TensorHandleResolve(_eagerTensorHandle, tf.Status.Handle); | ||
| 72 | 69 | tf.Status.Check(true); | |
| 73 | - #if TRACK_TENSOR_LIFE | ||
| 74 | - print($"Take EagerTensorHandle {EagerTensorHandle} {Id} Resolving 0x{_handle.ToString("x16")}"); | ||
| 75 | - #endif | ||
| 76 | 70 | } | |
| 77 | 71 | ||
| 78 | 72 | /// <summary> | |
@@ -104,7 +98,7 @@ void copy_handle_data(Tensor target_t) | |||
| 104 | 98 | protected override void DisposeUnmanagedResources(IntPtr handle) | |
| 105 | 99 | { | |
| 106 | 100 | base.DisposeUnmanagedResources(handle); | |
| 107 | - EagerTensorHandle.Dispose(); | ||
| 101 | + _eagerTensorHandle.Dispose(); | ||
| 108 | 102 | } | |
| 109 | 103 | } | |
| 110 | 104 | } | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -25,40 +25,62 @@ public static NDArray Scalar<T>(T value) where T : unmanaged | |||
| 25 | 25 | bool val => new NDArray(val), | |
| 26 | 26 | byte val => new NDArray(val), | |
| 27 | 27 | int val => new NDArray(val), | |
| 28 | + long val => new NDArray(val), | ||
| 28 | 29 | float val => new NDArray(val), | |
| 29 | 30 | double val => new NDArray(val), | |
| 30 | 31 | _ => throw new NotImplementedException("") | |
| 31 | 32 | }; | |
| 32 | 33 | ||
| 33 | 34 | void Init<T>(T value) where T : unmanaged | |
| 34 | 35 | { | |
| 35 | - _tensor = new EagerTensor(value); | ||
| 36 | + _tensor = value switch | ||
| 37 | + { | ||
| 38 | + bool val => new Tensor(val), | ||
| 39 | + byte val => new Tensor(val), | ||
| 40 | + int val => new Tensor(val), | ||
| 41 | + long val => new Tensor(val), | ||
| 42 | + float val => new Tensor(val), | ||
| 43 | + double val => new Tensor(val), | ||
| 44 | + _ => throw new NotImplementedException("") | ||
| 45 | + }; | ||
| 36 | 46 | _tensor.SetReferencedByNDArray(); | |
| 47 | + | ||
| 48 | + var _handle = c_api.TFE_NewTensorHandle(_tensor, tf.Status.Handle); | ||
| 49 | + _tensor.SetEagerTensorHandle(_handle); | ||
| 37 | 50 | } | |
| 38 | 51 | ||
| 39 | 52 | void Init(Array value, Shape? shape = null) | |
| 40 | 53 | { | |
| 41 | - _tensor = new EagerTensor(value, shape ?? value.GetShape()); | ||
| 54 | + _tensor = new Tensor(value, shape ?? value.GetShape()); | ||
| 42 | 55 | _tensor.SetReferencedByNDArray(); | |
| 56 | + | ||
| 57 | + var _handle = c_api.TFE_NewTensorHandle(_tensor, tf.Status.Handle); | ||
| 58 | + _tensor.SetEagerTensorHandle(_handle); | ||
| 43 | 59 | } | |
| 44 | 60 | ||
| 45 | 61 | void Init(Shape shape, TF_DataType dtype = TF_DataType.TF_DOUBLE) | |
| 46 | 62 | { | |
| 47 | - _tensor = new EagerTensor(shape, dtype: dtype); | ||
| 63 | + _tensor = new Tensor(shape, dtype: dtype); | ||
| 48 | 64 | _tensor.SetReferencedByNDArray(); | |
| 65 | + | ||
| 66 | + var _handle = c_api.TFE_NewTensorHandle(_tensor, tf.Status.Handle); | ||
| 67 | + _tensor.SetEagerTensorHandle(_handle); | ||
| 49 | 68 | } | |
| 50 | 69 | ||
| 51 | 70 | void Init(Tensor value, Shape? shape = null) | |
| 52 | 71 | { | |
| 53 | 72 | if (shape is not null) | |
| 54 | - _tensor = tf.reshape(value, shape); | ||
| 73 | + _tensor = new Tensor(value.TensorDataPointer, shape, value.dtype); | ||
| 55 | 74 | else | |
| 56 | 75 | _tensor = value; | |
| 57 | 76 | ||
| 58 | 77 | if (_tensor.TensorDataPointer == IntPtr.Zero) | |
| 59 | 78 | _tensor = tf.get_default_session().eval(_tensor); | |
| 60 | 79 | ||
| 61 | 80 | _tensor.SetReferencedByNDArray(); | |
| 81 | + | ||
| 82 | + var _handle = c_api.TFE_NewTensorHandle(_tensor, tf.Status.Handle); | ||
| 83 | + _tensor.SetEagerTensorHandle(_handle); | ||
| 62 | 84 | } | |
| 63 | 85 | } | |
| 64 | 86 | } | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -6,7 +6,7 @@ | |||
| 6 | 6 | ||
| 7 | 7 | namespace Tensorflow.NumPy | |
| 8 | 8 | { | |
| 9 | - public partial class NDArray | ||
| 9 | + public partial class NDArray : DisposableObject | ||
| 10 | 10 | { | |
| 11 | 11 | Tensor _tensor; | |
| 12 | 12 | public TF_DataType dtype => _tensor.dtype; | |
@@ -58,5 +58,11 @@ public override string ToString() | |||
| 58 | 58 | { | |
| 59 | 59 | return tensor_util.to_numpy_string(_tensor); | |
| 60 | 60 | } | |
| 61 | + | ||
| 62 | + protected override void DisposeUnmanagedResources(IntPtr handle) | ||
| 63 | + { | ||
| 64 | + _tensor.EagerTensorHandle.Dispose(); | ||
| 65 | + _tensor.Dispose(); | ||
| 66 | + } | ||
| 61 | 67 | } | |
| 62 | 68 | } | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -53,10 +53,9 @@ public Tensor(IntPtr handle) | |||
| 53 | 53 | /// <param name="data_ptr">Pointer to unmanaged, fixed or pinned memory which the caller owns</param> | |
| 54 | 54 | /// <param name="shape">Tensor shape</param> | |
| 55 | 55 | /// <param name="dType">TF data type</param> | |
| 56 | - /// <param name="num_bytes">Size of the tensor in memory</param> | ||
| 57 | - public Tensor(IntPtr data_ptr, long[] shape, TF_DataType dType, int num_bytes) | ||
| 56 | + public unsafe Tensor(IntPtr data_ptr, Shape shape, TF_DataType dtype) | ||
| 58 | 57 | { | |
| 59 | - _handle = TF_NewTensor(dType, dims: shape, num_dims: shape.Length, data: data_ptr, len: (ulong)num_bytes); | ||
| 58 | + _handle = TF_NewTensor(shape, dtype, data: data_ptr.ToPointer()); | ||
| 60 | 59 | isCreatedInGraphMode = !tf.executing_eagerly(); | |
| 61 | 60 | } | |
| 62 | 61 | ||
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -89,10 +89,11 @@ public partial class Tensor : DisposableObject, | |||
| 89 | 89 | /// </summary> | |
| 90 | 90 | public object Tag { get; set; } | |
| 91 | 91 | ||
| 92 | + protected SafeTensorHandleHandle _eagerTensorHandle; | ||
| 92 | 93 | /// <summary> | |
| 93 | 94 | /// TFE_TensorHandle | |
| 94 | 95 | /// </summary> | |
| 95 | - public SafeTensorHandleHandle EagerTensorHandle { get; set; } | ||
| 96 | + public SafeTensorHandleHandle EagerTensorHandle => _eagerTensorHandle; | ||
| 96 | 97 | ||
| 97 | 98 | protected bool isReferencedByNDArray; | |
| 98 | 99 | public bool IsReferencedByNDArray => isReferencedByNDArray; | |
@@ -212,6 +213,7 @@ public TF_Output _as_tf_output() | |||
| 212 | 213 | } | |
| 213 | 214 | ||
| 214 | 215 | public void SetReferencedByNDArray() => isReferencedByNDArray = true; | |
| 216 | + public void SetEagerTensorHandle(SafeTensorHandleHandle handle) => _eagerTensorHandle = handle; | ||
| 215 | 217 | ||
| 216 | 218 | public Tensor MaybeMove() | |
| 217 | 219 | { | |
@@ -254,30 +256,16 @@ public override string ToString() | |||
| 254 | 256 | } | |
| 255 | 257 | } | |
| 256 | 258 | ||
| 257 | - /// <summary> | ||
| 258 | - /// Dispose any managed resources. | ||
| 259 | - /// </summary> | ||
| 260 | - /// <remarks>Equivalent to what you would perform inside <see cref="DisposableObject.Dispose"/></remarks> | ||
| 261 | - protected override void DisposeManagedResources() | ||
| 262 | - { | ||
| 263 | - | ||
| 264 | - } | ||
| 265 | - | ||
| 266 | 259 | [SuppressMessage("ReSharper", "ConvertIfStatementToSwitchStatement")] | |
| 267 | 260 | protected override void DisposeUnmanagedResources(IntPtr handle) | |
| 268 | 261 | { | |
| 269 | - #if TRACK_TENSOR_LIFE | ||
| 270 | - print($"Delete Tensor 0x{handle.ToString("x16")} {AllocationType} Data: 0x{TensorDataPointer.ToString("x16")}"); | ||
| 271 | - #endif | ||
| 272 | 262 | if (dtype == TF_DataType.TF_STRING) | |
| 273 | 263 | { | |
| 274 | 264 | long size = 1; | |
| 275 | 265 | foreach (var s in TensorShape.dims) | |
| 276 | 266 | size *= s; | |
| 277 | 267 | var tstr = TensorDataPointer; | |
| 278 | - #if TRACK_TENSOR_LIFE | ||
| 279 | - print($"Delete TString 0x{handle.ToString("x16")} {AllocationType} Data: 0x{tstr.ToString("x16")}"); | ||
| 280 | - #endif | ||
| 268 | + | ||
| 281 | 269 | for (int i = 0; i < size; i++) | |
| 282 | 270 | { | |
| 283 | 271 | c_api.TF_StringDealloc(tstr); | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -101,7 +101,7 @@ public partial class c_api | |||
| 101 | 101 | [MethodImpl(MethodImplOptions.AggressiveInlining)] | |
| 102 | 102 | public static unsafe IntPtr TF_NewTensor(TF_DataType dataType, long[] dims, int num_dims, IntPtr data, ulong len) | |
| 103 | 103 | { | |
| 104 | - return c_api.TF_NewTensor(dataType, dims, num_dims, data, len, EmptyDeallocator, DeallocatorArgs.Empty); | ||
| 104 | + return TF_NewTensor(dataType, dims, num_dims, data, len, EmptyDeallocator, DeallocatorArgs.Empty); | ||
| 105 | 105 | } | |
| 106 | 106 | ||
| 107 | 107 | public static unsafe IntPtr TF_NewTensor(Shape shape, TF_DataType dtype, void* data) | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -101,7 +101,7 @@ public static Tensor _constant_impl(object value, | |||
| 101 | 101 | return op.outputs[0]; | |
| 102 | 102 | } | |
| 103 | 103 | ||
| 104 | - private static Tensor _eager_reshape(EagerTensor tensor, int[] shape, Context ctx) | ||
| 104 | + private static Tensor _eager_reshape(Tensor tensor, int[] shape, Context ctx) | ||
| 105 | 105 | { | |
| 106 | 106 | var attr_t = tensor.dtype.as_datatype_enum(); | |
| 107 | 107 | var dims_t = convert_to_eager_tensor(shape, ctx, dtypes.int32); | |
@@ -111,7 +111,7 @@ private static Tensor _eager_reshape(EagerTensor tensor, int[] shape, Context ct | |||
| 111 | 111 | return result[0]; | |
| 112 | 112 | } | |
| 113 | 113 | ||
| 114 | - private static Tensor _eager_fill(int[] dims, EagerTensor value, Context ctx) | ||
| 114 | + private static Tensor _eager_fill(int[] dims, Tensor value, Context ctx) | ||
| 115 | 115 | { | |
| 116 | 116 | var attr_t = value.dtype.as_datatype_enum(); | |
| 117 | 117 | var dims_t = convert_to_eager_tensor(dims, ctx, dtypes.int32); | |
@@ -121,7 +121,7 @@ private static Tensor _eager_fill(int[] dims, EagerTensor value, Context ctx) | |||
| 121 | 121 | return result[0]; | |
| 122 | 122 | } | |
| 123 | 123 | ||
| 124 | - private static EagerTensor convert_to_eager_tensor(object value, Context ctx, TF_DataType dtype = TF_DataType.DtInvalid) | ||
| 124 | + private static Tensor convert_to_eager_tensor(object value, Context ctx, TF_DataType dtype = TF_DataType.DtInvalid) | ||
| 125 | 125 | { | |
| 126 | 126 | ctx.ensure_initialized(); | |
| 127 | 127 | // convert data type | |
@@ -161,7 +161,7 @@ value is NDArray nd && | |||
| 161 | 161 | case EagerTensor val: | |
| 162 | 162 | return val; | |
| 163 | 163 | case NDArray val: | |
| 164 | - return (EagerTensor)val; | ||
| 164 | + return val; | ||
| 165 | 165 | case Shape val: | |
| 166 | 166 | return new EagerTensor(val.dims, new Shape(val.ndim)); | |
| 167 | 167 | case TensorShape val: | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -18,7 +18,7 @@ public unsafe void TensorFromFixed() | |||
| 18 | 18 | var span = new Span<float>(array, 100, 500); | |
| 19 | 19 | fixed (float* ptr = &MemoryMarshal.GetReference(span)) | |
| 20 | 20 | { | |
| 21 | - using (var t = new Tensor((IntPtr)ptr, new long[] { span.Length }, tf.float32, 4 * span.Length)) | ||
| 21 | + using (var t = new Tensor((IntPtr)ptr, new long[] { span.Length }, tf.float32)) | ||
| 22 | 22 | { | |
| 23 | 23 | Assert.IsFalse(t.IsDisposed); | |
| 24 | 24 | Assert.AreEqual(2000, (int)t.bytesize); | |
@@ -27,7 +27,7 @@ public unsafe void TensorFromFixed() | |||
| 27 | 27 | ||
| 28 | 28 | fixed (float* ptr = &array[0]) | |
| 29 | 29 | { | |
| 30 | - using (var t = new Tensor((IntPtr)ptr, new long[] { array.Length }, tf.float32, 4 * array.Length)) | ||
| 30 | + using (var t = new Tensor((IntPtr)ptr, new long[] { array.Length }, tf.float32)) | ||
| 31 | 31 | { | |
| 32 | 32 | Assert.IsFalse(t.IsDisposed); | |
| 33 | 33 | Assert.AreEqual(4000, (int)t.bytesize); | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -14,11 +14,6 @@ public void TestInit() | |||
| 14 | 14 | tf.Context.ensure_initialized(); | |
| 15 | 15 | } | |
| 16 | 16 | ||
| 17 | - [TestCleanup] | ||
| 18 | - public void TestClean() | ||
| 19 | - { | ||
| 20 | - } | ||
| 21 | - | ||
| 22 | 17 | public bool Equal(float[] f1, float[] f2) | |
| 23 | 18 | { | |
| 24 | 19 | bool ret = false; | |
| Back | FazBrowse Home | New Git URL |
0 commit comments