| FazBrowse GitHub Viewer | Trending | | Home |
| Tools: [Download Repo ZIP] [Original HTTPS Page] |
1 parent ed8ccec commit 44d203d
4 files changed
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -8,6 +8,7 @@ namespace Tensorflow.NumPy | |||
| 8 | 8 | { | |
| 9 | 9 | public partial class NDArray | |
| 10 | 10 | { | |
| 11 | + protected NDArray() { } | ||
| 11 | 12 | public NDArray(bool value) : base(value) => NewEagerTensorHandle(); | |
| 12 | 13 | public NDArray(byte value) : base(value) => NewEagerTensorHandle(); | |
| 13 | 14 | public NDArray(short value) : base(value) => NewEagerTensorHandle(); | |
@@ -57,6 +58,20 @@ public static NDArray Scalar<T>(T value) where T : unmanaged | |||
| 57 | 58 | _ => throw new NotImplementedException("") | |
| 58 | 59 | }; | |
| 59 | 60 | ||
| 61 | + /// <summary> | ||
| 62 | + /// Reuse the existing memory instead of copying it. | ||
| 63 | + /// </summary> | ||
| 64 | + /// <param name="data_ptr"></param> | ||
| 65 | + /// <param name="shape"></param> | ||
| 66 | + /// <param name="dtype"></param> | ||
| 67 | + /// <param name="deallocator"></param> | ||
| 68 | + protected void InitWithExistingMemory(IntPtr data_ptr, Shape shape, TF_DataType dtype, c_api.DeallocatorV2 deallocator) | ||
| 69 | + { | ||
| 70 | + _handle = c_api.TF_NewTensor(TF_DataType.TF_STRING, shape.dims, shape.ndim, data_ptr, (ulong)(shape.size * dtype.get_datatype_size()), deallocator, IntPtr.Zero); | ||
| 71 | + tensor_util.DangerousManuallySetTensorDType(_handle, dtype); | ||
| 72 | + NewEagerTensorHandle(); | ||
| 73 | + } | ||
| 74 | + | ||
| 60 | 75 | void NewEagerTensorHandle() | |
| 61 | 76 | { | |
| 62 | 77 | if (_handle is not null) | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -417,7 +417,7 @@ public static Tensor ones(Shape shape, TF_DataType dtype = TF_DataType.TF_FLOAT, | |||
| 417 | 417 | { | |
| 418 | 418 | TF_DataType.TF_DOUBLE => constant(1.0d), | |
| 419 | 419 | TF_DataType.TF_FLOAT => constant(1.0f), | |
| 420 | - _ => constant(1) | ||
| 420 | + _ => constant(1, dtype) | ||
| 421 | 421 | }; | |
| 422 | 422 | ||
| 423 | 423 | if (shape.ndim == 0) | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -71,7 +71,7 @@ public partial class c_api | |||
| 71 | 71 | /// <param name="deallocator_arg"></param> | |
| 72 | 72 | /// <returns></returns> | |
| 73 | 73 | [DllImport(TensorFlowLibName)] | |
| 74 | - public static extern SafeTensorHandle TF_NewTensor(TF_DataType dataType, long[] dims, int num_dims, IntPtr data, ulong len, Deallocator deallocator, IntPtr deallocator_arg); | ||
| 74 | + public static extern SafeTensorHandle TF_NewTensor(TF_DataType dataType, long[] dims, int num_dims, IntPtr data, ulong len, DeallocatorV2 deallocator, IntPtr deallocator_arg); | ||
| 75 | 75 | ||
| 76 | 76 | public static unsafe SafeTensorHandle TF_NewTensor(byte[] data, Shape shape, TF_DataType dtype) | |
| 77 | 77 | { | |
@@ -147,6 +147,15 @@ public static unsafe SafeTensorHandle TF_NewTensor<T>(T value) | |||
| 147 | 147 | [DllImport(TensorFlowLibName)] | |
| 148 | 148 | public static extern TF_DataType TF_TensorType(SafeTensorHandle tensor); | |
| 149 | 149 | ||
| 150 | + /// <summary> | ||
| 151 | + /// Set a new shape for the Tensor. Note that this API only works after tf2.11. | ||
| 152 | + /// </summary> | ||
| 153 | + /// <param name="tensor"></param> | ||
| 154 | + /// <param name="dims"></param> | ||
| 155 | + /// <param name="num_dims"></param> | ||
| 156 | + [DllImport(TensorFlowLibName)] | ||
| 157 | + public static extern void TF_SetShape(SafeTensorHandle tensor, long[] dims, int num_dims); | ||
| 158 | + | ||
| 150 | 159 | /// <summary> | |
| 151 | 160 | /// Return the size in bytes required to encode a string `len` bytes long into a | |
| 152 | 161 | /// TF_STRING tensor. | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -22,6 +22,7 @@ limitations under the License. | |||
| 22 | 22 | using Tensorflow.Eager; | |
| 23 | 23 | using Tensorflow.Graphs; | |
| 24 | 24 | using static Tensorflow.Binding; | |
| 25 | + using System.Diagnostics; | ||
| 25 | 26 | ||
| 26 | 27 | namespace Tensorflow | |
| 27 | 28 | { | |
@@ -649,5 +650,24 @@ public static ParsedSliceArgs ParseSlices(Tensor start, Tensor stop = null, Tens | |||
| 649 | 650 | NewAxisMask = new_axis_mask | |
| 650 | 651 | }; | |
| 651 | 652 | } | |
| 653 | + | ||
| 654 | + /// <summary> | ||
| 655 | + /// Warning: this method is an extremely dangerous method. It directly changes the dtype inside the tensor | ||
| 656 | + /// and security is not guaranteed at all. Currently this method is only used for some conditions to reuse | ||
| 657 | + /// the existing memory. Any other usage should be prevented. If you are sure you want to use it when | ||
| 658 | + /// developing tensorflow.net, please ask @Oceanic2018 or @AsakusaRinne first. | ||
| 659 | + /// </summary> | ||
| 660 | + /// <param name="handle"></param> | ||
| 661 | + /// <param name="dtype"></param> | ||
| 662 | + internal static unsafe void DangerousManuallySetTensorDType(SafeTensorHandle handle, TF_DataType dtype) | ||
| 663 | + { | ||
| 664 | + long tf_tensor_address = handle.DangerousGetHandle().ToInt64(); | ||
| 665 | + long interface_address = *(long*)(tf_tensor_address); | ||
| 666 | + long tensor_shape_address = interface_address + 8; | ||
| 667 | + long tensor_dtype_address = tensor_shape_address + 13; | ||
| 668 | + byte* dtype_pointer = (byte*)tensor_dtype_address; | ||
| 669 | + *dtype_pointer = (byte)dtype; | ||
| 670 | + Debug.Assert(c_api.TF_TensorType(handle) == dtype); | ||
| 671 | + } | ||
| 652 | 672 | } | |
| 653 | 673 | } | |
| Back | FazBrowse Home | New Git URL |
0 commit comments