| FazBrowse GitHub Viewer | Trending | | Home |
| Tools: [Download Repo ZIP] [Original HTTPS Page] |
1 parent 94601f5 commit c7ee230
5 files changed
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -30,5 +30,62 @@ static T Scalar<T>(long input) | |||
| 30 | 30 | TypeCode.Single => (T)Convert.ChangeType(input, TypeCode.Single), | |
| 31 | 31 | _ => throw new NotImplementedException("") | |
| 32 | 32 | }; | |
| 33 | + | ||
| 34 | + public static unsafe Array ToMultiDimArray<T>(NDArray nd) where T : unmanaged | ||
| 35 | + { | ||
| 36 | + var ret = Array.CreateInstance(typeof(T), nd.shape.as_int_list()); | ||
| 37 | + | ||
| 38 | + var addr = ret switch | ||
| 39 | + { | ||
| 40 | + T[] array => Addr(array), | ||
| 41 | + T[,] array => Addr(array), | ||
| 42 | + T[,,] array => Addr(array), | ||
| 43 | + T[,,,] array => Addr(array), | ||
| 44 | + T[,,,,] array => Addr(array), | ||
| 45 | + T[,,,,,] array => Addr(array), | ||
| 46 | + _ => throw new NotImplementedException("") | ||
| 47 | + }; | ||
| 48 | + | ||
| 49 | + System.Buffer.MemoryCopy(nd.data.ToPointer(), addr, nd.bytesize, nd.bytesize); | ||
| 50 | + return ret; | ||
| 51 | + } | ||
| 52 | + | ||
| 53 | + #region multiple array | ||
| 54 | + static unsafe T* Addr<T>(T[] array) where T : unmanaged | ||
| 55 | + { | ||
| 56 | + fixed (T* a = &array[0]) | ||
| 57 | + return a; | ||
| 58 | + } | ||
| 59 | + | ||
| 60 | + static unsafe T* Addr<T>(T[,] array) where T : unmanaged | ||
| 61 | + { | ||
| 62 | + fixed (T* a = &array[0, 0]) | ||
| 63 | + return a; | ||
| 64 | + } | ||
| 65 | + | ||
| 66 | + static unsafe T* Addr<T>(T[,,] array) where T : unmanaged | ||
| 67 | + { | ||
| 68 | + fixed (T* a = &array[0, 0, 0]) | ||
| 69 | + return a; | ||
| 70 | + } | ||
| 71 | + | ||
| 72 | + static unsafe T* Addr<T>(T[,,,] array) where T : unmanaged | ||
| 73 | + { | ||
| 74 | + fixed (T* a = &array[0, 0, 0, 0]) | ||
| 75 | + return a; | ||
| 76 | + } | ||
| 77 | + | ||
| 78 | + static unsafe T* Addr<T>(T[,,,,] array) where T : unmanaged | ||
| 79 | + { | ||
| 80 | + fixed (T* a = &array[0, 0, 0, 0, 0]) | ||
| 81 | + return a; | ||
| 82 | + } | ||
| 83 | + | ||
| 84 | + static unsafe T* Addr<T>(T[,,,,,] array) where T : unmanaged | ||
| 85 | + { | ||
| 86 | + fixed (T* a = &array[0, 0, 0, 0, 0, 0]) | ||
| 87 | + return a; | ||
| 88 | + } | ||
| 89 | + #endregion | ||
| 33 | 90 | } | |
| 34 | 91 | } | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -8,28 +8,28 @@ namespace Tensorflow.NumPy | |||
| 8 | 8 | { | |
| 9 | 9 | public partial class NDArray | |
| 10 | 10 | { | |
| 11 | - public NDArray(bool value) : base(value) { NewEagerTensorHandle(); } | ||
| 12 | - public NDArray(byte value) : base(value) { NewEagerTensorHandle(); } | ||
| 13 | - public NDArray(short value) : base(value) { NewEagerTensorHandle(); } | ||
| 14 | - public NDArray(int value) : base(value) { NewEagerTensorHandle(); } | ||
| 15 | - public NDArray(long value) : base(value) { NewEagerTensorHandle(); } | ||
| 16 | - public NDArray(float value) : base(value) { NewEagerTensorHandle(); } | ||
| 17 | - public NDArray(double value) : base(value) { NewEagerTensorHandle(); } | ||
| 11 | + public NDArray(bool value) : base(value) => NewEagerTensorHandle(); | ||
| 12 | + public NDArray(byte value) : base(value) => NewEagerTensorHandle(); | ||
| 13 | + public NDArray(short value) : base(value) => NewEagerTensorHandle(); | ||
| 14 | + public NDArray(int value) : base(value) => NewEagerTensorHandle(); | ||
| 15 | + public NDArray(long value) : base(value) => NewEagerTensorHandle(); | ||
| 16 | + public NDArray(float value) : base(value) => NewEagerTensorHandle(); | ||
| 17 | + public NDArray(double value) : base(value) => NewEagerTensorHandle(); | ||
| 18 | 18 | ||
| 19 | - public NDArray(Array value, Shape? shape = null) | ||
| 20 | - : base(value, shape) { NewEagerTensorHandle(); } | ||
| 19 | + public NDArray(Array value, Shape? shape = null) : base(value, shape) | ||
| 20 | + => NewEagerTensorHandle(); | ||
| 21 | 21 | ||
| 22 | - public NDArray(Shape shape, TF_DataType dtype = TF_DataType.TF_DOUBLE) | ||
| 23 | - : base(shape, dtype: dtype) { NewEagerTensorHandle(); } | ||
| 22 | + public NDArray(Shape shape, TF_DataType dtype = TF_DataType.TF_DOUBLE) : base(shape, dtype: dtype) | ||
| 23 | + => NewEagerTensorHandle(); | ||
| 24 | 24 | ||
| 25 | - public NDArray(byte[] bytes, Shape shape, TF_DataType dtype) | ||
| 26 | - : base(bytes, shape, dtype) { NewEagerTensorHandle(); } | ||
| 25 | + public NDArray(byte[] bytes, Shape shape, TF_DataType dtype) : base(bytes, shape, dtype) | ||
| 26 | + => NewEagerTensorHandle(); | ||
| 27 | 27 | ||
| 28 | - public NDArray(long[] value, Shape? shape = null) | ||
| 29 | - : base(value, shape) { NewEagerTensorHandle(); } | ||
| 28 | + public NDArray(long[] value, Shape? shape = null) : base(value, shape) | ||
| 29 | + => NewEagerTensorHandle(); | ||
| 30 | 30 | ||
| 31 | - public NDArray(IntPtr address, Shape shape, TF_DataType dtype) | ||
| 32 | - : base(address, shape, dtype) { NewEagerTensorHandle(); } | ||
| 31 | + public NDArray(IntPtr address, Shape shape, TF_DataType dtype) : base(address, shape, dtype) | ||
| 32 | + => NewEagerTensorHandle(); | ||
| 33 | 33 | ||
| 34 | 34 | public NDArray(Tensor tensor, bool clone = false) : base(tensor.Handle, clone: clone) | |
| 35 | 35 | { | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -19,6 +19,7 @@ limitations under the License. | |||
| 19 | 19 | using System.Collections.Generic; | |
| 20 | 20 | using System.Linq; | |
| 21 | 21 | using System.Text; | |
| 22 | + using Tensorflow.Util; | ||
| 22 | 23 | using static Tensorflow.Binding; | |
| 23 | 24 | ||
| 24 | 25 | namespace Tensorflow.NumPy | |
@@ -35,7 +36,10 @@ public ValueType GetValue(params int[] indices) | |||
| 35 | 36 | public NDArray astype(TF_DataType dtype) => new NDArray(math_ops.cast(this, dtype)); | |
| 36 | 37 | public NDArray ravel() => throw new NotImplementedException(""); | |
| 37 | 38 | public void shuffle(NDArray nd) => np.random.shuffle(nd); | |
| 38 | - public Array ToMuliDimArray<T>() => throw new NotImplementedException(""); | ||
| 39 | + | ||
| 40 | + public unsafe Array ToMultiDimArray<T>() where T : unmanaged | ||
| 41 | + => NDArrayConverter.ToMultiDimArray<T>(this); | ||
| 42 | + | ||
| 39 | 43 | public byte[] ToByteArray() => BufferToArray(); | |
| 40 | 44 | public override string ToString() => NDArrayRender.ToString(this); | |
| 41 | 45 | ||
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -273,19 +273,19 @@ private static void WriteDataset(long f, string name, Tensor data) | |||
| 273 | 273 | switch (data.dtype) | |
| 274 | 274 | { | |
| 275 | 275 | case TF_DataType.TF_FLOAT: | |
| 276 | - Hdf5.WriteDatasetFromArray<float>(f, name, data.numpy().ToMuliDimArray<float>()); | ||
| 276 | + Hdf5.WriteDatasetFromArray<float>(f, name, data.numpy().ToMultiDimArray<float>()); | ||
| 277 | 277 | break; | |
| 278 | 278 | case TF_DataType.TF_DOUBLE: | |
| 279 | - Hdf5.WriteDatasetFromArray<double>(f, name, data.numpy().ToMuliDimArray<double>()); | ||
| 279 | + Hdf5.WriteDatasetFromArray<double>(f, name, data.numpy().ToMultiDimArray<double>()); | ||
| 280 | 280 | break; | |
| 281 | 281 | case TF_DataType.TF_INT32: | |
| 282 | - Hdf5.WriteDatasetFromArray<int>(f, name, data.numpy().ToMuliDimArray<int>()); | ||
| 282 | + Hdf5.WriteDatasetFromArray<int>(f, name, data.numpy().ToMultiDimArray<int>()); | ||
| 283 | 283 | break; | |
| 284 | 284 | case TF_DataType.TF_INT64: | |
| 285 | - Hdf5.WriteDatasetFromArray<long>(f, name, data.numpy().ToMuliDimArray<long>()); | ||
| 285 | + Hdf5.WriteDatasetFromArray<long>(f, name, data.numpy().ToMultiDimArray<long>()); | ||
| 286 | 286 | break; | |
| 287 | 287 | default: | |
| 288 | - Hdf5.WriteDatasetFromArray<float>(f, name, data.numpy().ToMuliDimArray<float>()); | ||
| 288 | + Hdf5.WriteDatasetFromArray<float>(f, name, data.numpy().ToMultiDimArray<float>()); | ||
| 289 | 289 | break; | |
| 290 | 290 | } | |
| 291 | 291 | } | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -50,6 +50,22 @@ public void array() | |||
| 50 | 50 | AssetSequenceEqual(new[] { 1, 2, 3, 4, 5, 6 }, x.ToArray<int>()); | |
| 51 | 51 | } | |
| 52 | 52 | ||
| 53 | + [TestMethod] | ||
| 54 | + public void to_multi_dim_array() | ||
| 55 | + { | ||
| 56 | + var x1 = np.arange(12); | ||
| 57 | + var y1 = x1.ToMultiDimArray<int>(); | ||
| 58 | + AssetSequenceEqual((int[])y1, x1.ToArray<int>()); | ||
| 59 | + | ||
| 60 | + var x2 = np.arange(12).reshape((2, 6)); | ||
| 61 | + var y2 = (int[,])x2.ToMultiDimArray<int>(); | ||
| 62 | + Assert.AreEqual(x2[0, 5], y2[0, 5]); | ||
| 63 | + | ||
| 64 | + var x3 = np.arange(12).reshape((2, 2, 3)); | ||
| 65 | + var y3 = (int[,,])x3.ToMultiDimArray<int>(); | ||
| 66 | + Assert.AreEqual(x3[0, 1, 2], y3[0, 1, 2]); | ||
| 67 | + } | ||
| 68 | + | ||
| 53 | 69 | [TestMethod] | |
| 54 | 70 | public void eye() | |
| 55 | 71 | { | |
| Back | FazBrowse Home | New Git URL |
0 commit comments