| FazBrowse GitHub Viewer | Trending | | Home |
| Tools: [Download Repo ZIP] [Original HTTPS Page] |
1 parent e748a8f commit 6adcfae
8 files changed
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -60,6 +60,9 @@ void NewEagerTensorHandle(SafeTensorHandle h) | |||
| 60 | 60 | { | |
| 61 | 61 | _id = ops.uid(); | |
| 62 | 62 | _eagerTensorHandle = c_api.TFE_NewTensorHandle(h, tf.Status.Handle); | |
| 63 | + #if TRACK_TENSOR_LIFE | ||
| 64 | + Console.WriteLine($"New EagerTensor {_eagerTensorHandle}"); | ||
| 65 | + #endif | ||
| 63 | 66 | tf.Status.Check(true); | |
| 64 | 67 | } | |
| 65 | 68 | ||
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -16,6 +16,7 @@ public override bool Equals(object obj) | |||
| 16 | 16 | long val => GetAtIndex<long>(0) == val, | |
| 17 | 17 | float val => GetAtIndex<float>(0) == val, | |
| 18 | 18 | double val => GetAtIndex<double>(0) == val, | |
| 19 | + string val => StringData(0) == val, | ||
| 19 | 20 | NDArray val => Equals(this, val), | |
| 20 | 21 | _ => base.Equals(obj) | |
| 21 | 22 | }; | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -91,7 +91,7 @@ static string Render(NDArray array) | |||
| 91 | 91 | .Take(25) | |
| 92 | 92 | .Select(x => x < 32 || x > 127 ? "\\x" + x.ToString("x") : Convert.ToChar(x).ToString())) + "'"; | |
| 93 | 93 | else | |
| 94 | - return $"['{string.Join("', '", array.StringData().Take(25))}']"; | ||
| 94 | + return $"'{string.Join("', '", array.StringData().Take(25))}'"; | ||
| 95 | 95 | } | |
| 96 | 96 | else if (dtype == TF_DataType.TF_VARIANT) | |
| 97 | 97 | { | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -51,7 +51,7 @@ tf.net 0.6x.x aligns with TensorFlow v2.6.x native library.</PackageReleaseNotes | |||
| 51 | 51 | ||
| 52 | 52 | <PropertyGroup Condition="'$(Configuration)|$(Platform)'=='Debug|x64'"> | |
| 53 | 53 | <AllowUnsafeBlocks>true</AllowUnsafeBlocks> | |
| 54 | - <DefineConstants>TRACE;DEBUG</DefineConstants> | ||
| 54 | + <DefineConstants>TRACE;DEBUG;TRACK_TENSOR_LIFE1</DefineConstants> | ||
| 55 | 55 | <PlatformTarget>x64</PlatformTarget> | |
| 56 | 56 | <DocumentationFile>TensorFlow.NET.xml</DocumentationFile> | |
| 57 | 57 | </PropertyGroup> | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -8,7 +8,7 @@ namespace Tensorflow | |||
| 8 | 8 | public sealed class SafeStringTensorHandle : SafeTensorHandle | |
| 9 | 9 | { | |
| 10 | 10 | Shape _shape; | |
| 11 | - IntPtr _handle; | ||
| 11 | + SafeTensorHandle _tensorHandle; | ||
| 12 | 12 | const int TF_TSRING_SIZE = 24; | |
| 13 | 13 | ||
| 14 | 14 | protected SafeStringTensorHandle() | |
@@ -18,23 +18,26 @@ protected SafeStringTensorHandle() | |||
| 18 | 18 | public SafeStringTensorHandle(SafeTensorHandle handle, Shape shape) | |
| 19 | 19 | : base(handle.DangerousGetHandle()) | |
| 20 | 20 | { | |
| 21 | - _handle = c_api.TF_TensorData(handle); | ||
| 21 | + _tensorHandle = handle; | ||
| 22 | 22 | _shape = shape; | |
| 23 | + bool success = false; | ||
| 24 | + _tensorHandle.DangerousAddRef(ref success); | ||
| 23 | 25 | } | |
| 24 | 26 | ||
| 25 | 27 | protected override bool ReleaseHandle() | |
| 26 | 28 | { | |
| 29 | + var _handle = c_api.TF_TensorData(_tensorHandle); | ||
| 27 | 30 | #if TRACK_TENSOR_LIFE | |
| 28 | - print($"Delete StringTensorHandle 0x{handle.ToString("x16")}"); | ||
| 31 | + Console.WriteLine($"Delete StringTensorData 0x{_handle.ToString("x16")}"); | ||
| 29 | 32 | #endif | |
| 30 | - | ||
| 31 | 33 | for (int i = 0; i < _shape.size; i++) | |
| 32 | 34 | { | |
| 33 | 35 | c_api.TF_StringDealloc(_handle); | |
| 34 | 36 | _handle += TF_TSRING_SIZE; | |
| 35 | 37 | } | |
| 36 | 38 | ||
| 37 | 39 | SetHandle(IntPtr.Zero); | |
| 40 | + _tensorHandle.DangerousRelease(); | ||
| 38 | 41 | ||
| 39 | 42 | return true; | |
| 40 | 43 | } | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -29,13 +29,13 @@ public SafeStringTensorHandle StringTensor(byte[][] buffer, Shape shape) | |||
| 29 | 29 | ||
| 30 | 30 | var tstr = c_api.TF_TensorData(handle); | |
| 31 | 31 | #if TRACK_TENSOR_LIFE | |
| 32 | - print($"New TString 0x{handle.ToString("x16")} Data: 0x{tstr.ToString("x16")}"); | ||
| 32 | + print($"New StringTensor {handle} Data: 0x{tstr.ToString("x16")}"); | ||
| 33 | 33 | #endif | |
| 34 | 34 | for (int i = 0; i < buffer.Length; i++) | |
| 35 | 35 | { | |
| 36 | 36 | c_api.TF_StringInit(tstr); | |
| 37 | 37 | c_api.TF_StringCopy(tstr, buffer[i], buffer[i].Length); | |
| 38 | - var data = c_api.TF_StringGetDataPointer(tstr); | ||
| 38 | + // var data = c_api.TF_StringGetDataPointer(tstr); | ||
| 39 | 39 | tstr += TF_TSRING_SIZE; | |
| 40 | 40 | } | |
| 41 | 41 | ||
@@ -53,6 +53,36 @@ public string[] StringData() | |||
| 53 | 53 | return _str; | |
| 54 | 54 | } | |
| 55 | 55 | ||
| 56 | + public string StringData(int index) | ||
| 57 | + { | ||
| 58 | + var bytes = StringBytes(index); | ||
| 59 | + return Encoding.UTF8.GetString(bytes); | ||
| 60 | + } | ||
| 61 | + | ||
| 62 | + public byte[] StringBytes(int index) | ||
| 63 | + { | ||
| 64 | + if (dtype != TF_DataType.TF_STRING) | ||
| 65 | + throw new InvalidOperationException($"Unable to call StringData when dtype != TF_DataType.TF_STRING (dtype is {dtype})"); | ||
| 66 | + | ||
| 67 | + byte[] buffer = new byte[0]; | ||
| 68 | + var tstrings = TensorDataPointer; | ||
| 69 | + for (int i = 0; i < shape.size; i++) | ||
| 70 | + { | ||
| 71 | + if(index == i) | ||
| 72 | + { | ||
| 73 | + var data = c_api.TF_StringGetDataPointer(tstrings); | ||
| 74 | + var len = c_api.TF_StringGetSize(tstrings); | ||
| 75 | + buffer = new byte[len]; | ||
| 76 | + // var capacity = c_api.TF_StringGetCapacity(tstrings); | ||
| 77 | + // var type = c_api.TF_StringGetType(tstrings); | ||
| 78 | + Marshal.Copy(data, buffer, 0, Convert.ToInt32(len)); | ||
| 79 | + break; | ||
| 80 | + } | ||
| 81 | + tstrings += TF_TSRING_SIZE; | ||
| 82 | + } | ||
| 83 | + return buffer; | ||
| 84 | + } | ||
| 85 | + | ||
| 56 | 86 | public byte[][] StringBytes() | |
| 57 | 87 | { | |
| 58 | 88 | if (dtype != TF_DataType.TF_STRING) | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -22,7 +22,7 @@ public static NDArray to_categorical(NDArray y, int num_classes = -1, TF_DataTyp | |||
| 22 | 22 | // categorical[np.arange(y.size), y] = 1; | |
| 23 | 23 | for (var i = 0; i < (int)y.size; i++) | |
| 24 | 24 | { | |
| 25 | - categorical[i][y1[i]] = 1.0f; | ||
| 25 | + categorical[i, y1[i]] = 1.0f; | ||
| 26 | 26 | } | |
| 27 | 27 | ||
| 28 | 28 | return categorical; | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -51,12 +51,11 @@ public void StringArray() | |||
| 51 | 51 | { | |
| 52 | 52 | var strings = new[] { "map_and_batch_fusion", "noop_elimination", "shuffle_and_repeat_fusion" }; | |
| 53 | 53 | var tensor = tf.constant(strings, dtype: tf.@string, name: "optimizations"); | |
| 54 | - var stringData = tensor.StringData(); | ||
| 55 | 54 | ||
| 56 | 55 | Assert.AreEqual(3, tensor.shape[0]); | |
| 57 | - Assert.AreEqual(strings[0], stringData[0]); | ||
| 58 | - Assert.AreEqual(strings[1], stringData[1]); | ||
| 59 | - Assert.AreEqual(strings[2], stringData[2]); | ||
| 56 | + Assert.AreEqual(tensor[0].numpy(), strings[0]); | ||
| 57 | + Assert.AreEqual(tensor[1].numpy(), strings[1]); | ||
| 58 | + Assert.AreEqual(tensor[2].numpy(), strings[2]); | ||
| 60 | 59 | } | |
| 61 | 60 | ||
| 62 | 61 | [TestMethod] | |
| Back | FazBrowse Home | New Git URL |
0 commit comments