| FazBrowse GitHub Viewer | Trending | | Home |
| Tools: [Download Repo ZIP] [Original HTTPS Page] |
1 parent 5720dfd commit 8ae2feb
8 files changed
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -76,12 +76,12 @@ private ITensorOrOperation _as_graph_element_locked(object obj, bool allow_tenso | |||
| 76 | 76 | obj = temp_obj; | |
| 77 | 77 | ||
| 78 | 78 | // If obj appears to be a name... | |
| 79 | - if (obj is String str) | ||
| 79 | + if (obj is string name) | ||
| 80 | 80 | { | |
| 81 | - if(str.Contains(":") && allow_tensor) | ||
| 81 | + if(name.Contains(":") && allow_tensor) | ||
| 82 | 82 | { | |
| 83 | - string op_name = str.Split(':')[0]; | ||
| 84 | - int out_n = int.Parse(str.Split(':')[1]); | ||
| 83 | + string op_name = name.Split(':')[0]; | ||
| 84 | + int out_n = int.Parse(name.Split(':')[1]); | ||
| 85 | 85 | ||
| 86 | 86 | if (_nodes_by_name.ContainsKey(op_name)) | |
| 87 | 87 | return _nodes_by_name[op_name].outputs[out_n]; | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -67,7 +67,7 @@ private NDArray _run(object fetches, FeedItem[] feed_dict = null) | |||
| 67 | 67 | default: | |
| 68 | 68 | throw new NotImplementedException("_run subfeed"); | |
| 69 | 69 | } | |
| 70 | - feed_map[subfeed_t.name] = new Tuple<object, object>(subfeed_t, subfeed.Value); | ||
| 70 | + feed_map[subfeed_t.name] = (subfeed_t, subfeed.Value); | ||
| 71 | 71 | } | |
| 72 | 72 | } | |
| 73 | 73 | ||
@@ -178,7 +178,8 @@ private unsafe NDArray fetchValue(IntPtr output) | |||
| 178 | 178 | case TF_DataType.TF_STRING: | |
| 179 | 179 | var bytes = tensor.Data(); | |
| 180 | 180 | // wired, don't know why we have to start from offset 9. | |
| 181 | - var str = UTF8Encoding.Default.GetString(bytes, 9, bytes.Length - 9); | ||
| 181 | + // length in the begin | ||
| 182 | + var str = UTF8Encoding.Default.GetString(bytes, 9, bytes[8]); | ||
| 182 | 183 | nd = np.array(str).reshape(); | |
| 183 | 184 | break; | |
| 184 | 185 | case TF_DataType.TF_INT16: | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -0,0 +1,111 @@ | |||
| 1 | + using NumSharp.Core; | ||
| 2 | + using System; | ||
| 3 | + using System.Collections.Generic; | ||
| 4 | + using System.Linq; | ||
| 5 | + using System.Runtime.InteropServices; | ||
| 6 | + using System.Text; | ||
| 7 | + using static Tensorflow.c_api; | ||
| 8 | + | ||
| 9 | + namespace Tensorflow | ||
| 10 | + { | ||
| 11 | + public partial class Tensor | ||
| 12 | + { | ||
| 13 | + /// <summary> | ||
| 14 | + /// if original buffer is free. | ||
| 15 | + /// </summary> | ||
| 16 | + private bool deallocator_called; | ||
| 17 | + | ||
| 18 | + public Tensor(IntPtr handle) | ||
| 19 | + { | ||
| 20 | + _handle = handle; | ||
| 21 | + } | ||
| 22 | + | ||
| 23 | + public Tensor(NDArray nd) | ||
| 24 | + { | ||
| 25 | + _handle = Allocate(nd); | ||
| 26 | + } | ||
| 27 | + | ||
| 28 | + private IntPtr Allocate(NDArray nd) | ||
| 29 | + { | ||
| 30 | + IntPtr dotHandle = IntPtr.Zero; | ||
| 31 | + ulong size = 0; | ||
| 32 | + | ||
| 33 | + if (nd.dtype.Name != "String") | ||
| 34 | + { | ||
| 35 | + dotHandle = Marshal.AllocHGlobal(nd.dtypesize * nd.size); | ||
| 36 | + size = (ulong)(nd.size * nd.dtypesize); | ||
| 37 | + } | ||
| 38 | + | ||
| 39 | + switch (nd.dtype.Name) | ||
| 40 | + { | ||
| 41 | + case "Int16": | ||
| 42 | + Marshal.Copy(nd.Data<short>(), 0, dotHandle, nd.size); | ||
| 43 | + break; | ||
| 44 | + case "Int32": | ||
| 45 | + Marshal.Copy(nd.Data<int>(), 0, dotHandle, nd.size); | ||
| 46 | + break; | ||
| 47 | + case "Single": | ||
| 48 | + Marshal.Copy(nd.Data<float>(), 0, dotHandle, nd.size); | ||
| 49 | + break; | ||
| 50 | + case "Double": | ||
| 51 | + Marshal.Copy(nd.Data<double>(), 0, dotHandle, nd.size); | ||
| 52 | + break; | ||
| 53 | + case "String": | ||
| 54 | + /*var value = nd.Data<string>()[0]; | ||
| 55 | + var bytes = Encoding.UTF8.GetBytes(value); | ||
| 56 | + dotHandle = Marshal.AllocHGlobal(bytes.Length + 1); | ||
| 57 | + Marshal.Copy(bytes, 0, dotHandle, bytes.Length); | ||
| 58 | + size = (ulong)bytes.Length;*/ | ||
| 59 | + | ||
| 60 | + var str = nd.Data<string>()[0]; | ||
| 61 | + ulong dst_len = c_api.TF_StringEncodedSize((ulong)str.Length); | ||
| 62 | + //dotHandle = Marshal.AllocHGlobal((int)dst_len); | ||
| 63 | + //size = c_api.TF_StringEncode(str, (ulong)str.Length, dotHandle, dst_len, status); | ||
| 64 | + | ||
| 65 | + var dataType1 = ToTFDataType(nd.dtype); | ||
| 66 | + // shape | ||
| 67 | + var dims1 = nd.shape.Select(x => (long)x).ToArray(); | ||
| 68 | + | ||
| 69 | + var tfHandle1 = c_api.TF_AllocateTensor(dataType1, | ||
| 70 | + dims1, | ||
| 71 | + nd.ndim, | ||
| 72 | + dst_len); | ||
| 73 | + | ||
| 74 | + dotHandle = c_api.TF_TensorData(tfHandle1); | ||
| 75 | + c_api.TF_StringEncode(str, (ulong)str.Length, dotHandle, dst_len, status); | ||
| 76 | + return tfHandle1; | ||
| 77 | + break; | ||
| 78 | + default: | ||
| 79 | + throw new NotImplementedException("Marshal.Copy failed."); | ||
| 80 | + } | ||
| 81 | + | ||
| 82 | + var dataType = ToTFDataType(nd.dtype); | ||
| 83 | + // shape | ||
| 84 | + var dims = nd.shape.Select(x => (long)x).ToArray(); | ||
| 85 | + // Free the original buffer and set flag | ||
| 86 | + Deallocator deallocator = (IntPtr values, IntPtr len, ref bool closure) => | ||
| 87 | + { | ||
| 88 | + Marshal.FreeHGlobal(dotHandle); | ||
| 89 | + closure = true; | ||
| 90 | + }; | ||
| 91 | + | ||
| 92 | + var tfHandle = c_api.TF_NewTensor(dataType, | ||
| 93 | + dims, | ||
| 94 | + nd.ndim, | ||
| 95 | + dotHandle, | ||
| 96 | + size, | ||
| 97 | + deallocator, | ||
| 98 | + ref deallocator_called); | ||
| 99 | + | ||
| 100 | + return tfHandle; | ||
| 101 | + } | ||
| 102 | + | ||
| 103 | + public Tensor(Operation op, int value_index, TF_DataType dtype) | ||
| 104 | + { | ||
| 105 | + this.op = op; | ||
| 106 | + this.value_index = value_index; | ||
| 107 | + this._dtype = dtype; | ||
| 108 | + _id = ops.uid(); | ||
| 109 | + } | ||
| 110 | + } | ||
| 111 | + } | ||
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -95,86 +95,6 @@ public int rank | |||
| 95 | 95 | ||
| 96 | 96 | public int NDims => rank; | |
| 97 | 97 | ||
| 98 | - /// <summary> | ||
| 99 | - /// if original buffer is free. | ||
| 100 | - /// </summary> | ||
| 101 | - private bool deallocator_called; | ||
| 102 | - | ||
| 103 | - public Tensor(IntPtr handle) | ||
| 104 | - { | ||
| 105 | - _handle = handle; | ||
| 106 | - } | ||
| 107 | - | ||
| 108 | - public Tensor(NDArray nd) | ||
| 109 | - { | ||
| 110 | - _handle = Allocate(nd); | ||
| 111 | - } | ||
| 112 | - | ||
| 113 | - private IntPtr Allocate(NDArray nd) | ||
| 114 | - { | ||
| 115 | - IntPtr dotHandle = IntPtr.Zero; | ||
| 116 | - ulong size = 0; | ||
| 117 | - | ||
| 118 | - if (nd.dtype.Name != "String") | ||
| 119 | - { | ||
| 120 | - dotHandle = Marshal.AllocHGlobal(nd.dtypesize * nd.size); | ||
| 121 | - size = (ulong)(nd.size * nd.dtypesize); | ||
| 122 | - } | ||
| 123 | - | ||
| 124 | - switch (nd.dtype.Name) | ||
| 125 | - { | ||
| 126 | - case "Int16": | ||
| 127 | - Marshal.Copy(nd.Data<short>(), 0, dotHandle, nd.size); | ||
| 128 | - break; | ||
| 129 | - case "Int32": | ||
| 130 | - Marshal.Copy(nd.Data<int>(), 0, dotHandle, nd.size); | ||
| 131 | - break; | ||
| 132 | - case "Single": | ||
| 133 | - Marshal.Copy(nd.Data<float>(), 0, dotHandle, nd.size); | ||
| 134 | - break; | ||
| 135 | - case "Double": | ||
| 136 | - Marshal.Copy(nd.Data<double>(), 0, dotHandle, nd.size); | ||
| 137 | - break; | ||
| 138 | - case "String": | ||
| 139 | - var value = nd.Data<string>()[0]; | ||
| 140 | - var bytes = Encoding.UTF8.GetBytes(value); | ||
| 141 | - dotHandle = Marshal.AllocHGlobal(bytes.Length + 1); | ||
| 142 | - Marshal.Copy(bytes, 0, dotHandle, bytes.Length); | ||
| 143 | - size = (ulong)bytes.Length; | ||
| 144 | - break; | ||
| 145 | - default: | ||
| 146 | - throw new NotImplementedException("Marshal.Copy failed."); | ||
| 147 | - } | ||
| 148 | - | ||
| 149 | - var dataType = ToTFDataType(nd.dtype); | ||
| 150 | - // shape | ||
| 151 | - var dims = nd.shape.Select(x => (long)x).ToArray(); | ||
| 152 | - // Free the original buffer and set flag | ||
| 153 | - Deallocator deallocator = (IntPtr values, IntPtr len, ref bool closure) => | ||
| 154 | - { | ||
| 155 | - Marshal.FreeHGlobal(dotHandle); | ||
| 156 | - closure = true; | ||
| 157 | - }; | ||
| 158 | - | ||
| 159 | - var tfHandle = c_api.TF_NewTensor(dataType, | ||
| 160 | - dims, | ||
| 161 | - nd.ndim, | ||
| 162 | - dotHandle, | ||
| 163 | - size, | ||
| 164 | - deallocator, | ||
| 165 | - ref deallocator_called); | ||
| 166 | - | ||
| 167 | - return tfHandle; | ||
| 168 | - } | ||
| 169 | - | ||
| 170 | - public Tensor(Operation op, int value_index, TF_DataType dtype) | ||
| 171 | - { | ||
| 172 | - this.op = op; | ||
| 173 | - this.value_index = value_index; | ||
| 174 | - this._dtype = dtype; | ||
| 175 | - _id = ops.uid(); | ||
| 176 | - } | ||
| 177 | - | ||
| 178 | 98 | public Operation[] Consumers => consumers(); | |
| 179 | 99 | ||
| 180 | 100 | public string Device => op.Device; | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -120,7 +120,7 @@ public partial class c_api | |||
| 120 | 120 | /// <param name="status">TF_Status*</param> | |
| 121 | 121 | /// <returns>On success returns the size in bytes of the encoded string.</returns> | |
| 122 | 122 | [DllImport(TensorFlowLibName)] | |
| 123 | - public static extern ulong TF_StringEncode(string src, ulong src_len, string dst, ulong dst_len, IntPtr status); | ||
| 123 | + public static extern ulong TF_StringEncode(string src, ulong src_len, IntPtr dst, ulong dst_len, IntPtr status); | ||
| 124 | 124 | ||
| 125 | 125 | /// <summary> | |
| 126 | 126 | /// Decode a string encoded using TF_StringEncode. | |
@@ -132,6 +132,6 @@ public partial class c_api | |||
| 132 | 132 | /// <param name="status">TF_Status*</param> | |
| 133 | 133 | /// <returns></returns> | |
| 134 | 134 | [DllImport(TensorFlowLibName)] | |
| 135 | - public static extern ulong TF_StringDecode(string src, ulong src_len, IntPtr dst, ref ulong dst_len, IntPtr status); | ||
| 135 | + public static extern ulong TF_StringDecode(IntPtr src, ulong src_len, IntPtr dst, ref ulong dst_len, IntPtr status); | ||
| 136 | 136 | } | |
| 137 | 137 | } | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -138,12 +138,14 @@ private void _check_saver_def() | |||
| 138 | 138 | public string save(Session sess, | |
| 139 | 139 | string save_path, | |
| 140 | 140 | string global_step = "", | |
| 141 | + string latest_filename = "", | ||
| 141 | 142 | string meta_graph_suffix = "meta", | |
| 142 | 143 | bool write_meta_graph = true, | |
| 143 | 144 | bool write_state = true, | |
| 144 | 145 | bool strip_default_attrs = false) | |
| 145 | 146 | { | |
| 146 | - string latest_filename = "checkpoint"; | ||
| 147 | + if (string.IsNullOrEmpty(latest_filename)) | ||
| 148 | + latest_filename = "checkpoint"; | ||
| 147 | 149 | string model_checkpoint_path = ""; | |
| 148 | 150 | string checkpoint_file = ""; | |
| 149 | 151 | ||
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -3,6 +3,7 @@ | |||
| 3 | 3 | using System; | |
| 4 | 4 | using System.Collections.Generic; | |
| 5 | 5 | using System.Linq; | |
| 6 | + using System.Runtime.InteropServices; | ||
| 6 | 7 | using System.Text; | |
| 7 | 8 | using Tensorflow; | |
| 8 | 9 | ||
@@ -104,11 +105,14 @@ public void StringEncode() | |||
| 104 | 105 | string str = "Hello, TensorFlow.NET!"; | |
| 105 | 106 | ulong dst_len = c_api.TF_StringEncodedSize((ulong)str.Length); | |
| 106 | 107 | Assert.AreEqual(dst_len, (ulong)23); | |
| 107 | - string dst = ""; | ||
| 108 | - c_api.TF_StringEncode(str, (ulong)str.Length, dst, dst_len, status); | ||
| 108 | + IntPtr dst = Marshal.AllocHGlobal((int)dst_len); | ||
| 109 | + ulong encoded_len = c_api.TF_StringEncode(str, (ulong)str.Length, dst, dst_len, status); | ||
| 110 | + Assert.AreEqual((ulong)23, encoded_len); | ||
| 109 | 111 | Assert.AreEqual(status.Code, TF_Code.TF_OK); | |
| 110 | - | ||
| 111 | - //c_api.TF_StringDecode(str, (ulong)str.Length, IntPtr.Zero, ref dst_len, status); | ||
| 112 | + string encoded_str = Marshal.PtrToStringUTF8(dst + sizeof(byte)); | ||
| 113 | + Assert.AreEqual(encoded_str, str); | ||
| 114 | + Assert.AreEqual(str.Length, Marshal.ReadByte(dst)); | ||
| 115 | + //c_api.TF_StringDecode(dst, (ulong)str.Length, IntPtr.Zero, ref dst_len, status); | ||
| 112 | 116 | } | |
| 113 | 117 | ||
| 114 | 118 | /// <summary> | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -45,7 +45,6 @@ public void Save1() | |||
| 45 | 45 | }); | |
| 46 | 46 | } | |
| 47 | 47 | ||
| 48 | - [TestMethod] | ||
| 49 | 48 | public void Save2() | |
| 50 | 49 | { | |
| 51 | 50 | var v1 = tf.get_variable("v1", shape: new TensorShape(3), initializer: tf.zeros_initializer); | |
| Back | FazBrowse Home | New Git URL |
0 commit comments