| FazBrowse GitHub Viewer | Trending | | Home |
| Tools: [Download Repo ZIP] [Original HTTPS Page] |
1 parent 8ae2feb commit 2579afc
2 files changed
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -36,6 +36,10 @@ private IntPtr Allocate(NDArray nd) | |||
| 36 | 36 | size = (ulong)(nd.size * nd.dtypesize); | |
| 37 | 37 | } | |
| 38 | 38 | ||
| 39 | + var dataType = ToTFDataType(nd.dtype); | ||
| 40 | + // shape | ||
| 41 | + var dims = nd.shape.Select(x => (long)x).ToArray(); | ||
| 42 | + | ||
| 39 | 43 | switch (nd.dtype.Name) | |
| 40 | 44 | { | |
| 41 | 45 | case "Int16": | |
@@ -51,37 +55,25 @@ private IntPtr Allocate(NDArray nd) | |||
| 51 | 55 | Marshal.Copy(nd.Data<double>(), 0, dotHandle, nd.size); | |
| 52 | 56 | break; | |
| 53 | 57 | 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 | 58 | var str = nd.Data<string>()[0]; | |
| 61 | 59 | 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 | 60 | var dataType1 = ToTFDataType(nd.dtype); | |
| 66 | 61 | // shape | |
| 67 | 62 | var dims1 = nd.shape.Select(x => (long)x).ToArray(); | |
| 68 | 63 | ||
| 69 | 64 | var tfHandle1 = c_api.TF_AllocateTensor(dataType1, | |
| 70 | 65 | dims1, | |
| 71 | 66 | nd.ndim, | |
| 72 | - dst_len); | ||
| 67 | + dst_len + sizeof(Int64)); | ||
| 73 | 68 | ||
| 74 | 69 | dotHandle = c_api.TF_TensorData(tfHandle1); | |
| 75 | - c_api.TF_StringEncode(str, (ulong)str.Length, dotHandle, dst_len, status); | ||
| 70 | + Marshal.WriteInt64(dotHandle, 0); | ||
| 71 | + c_api.TF_StringEncode(str, (ulong)str.Length, dotHandle + sizeof(Int64), dst_len, status); | ||
| 76 | 72 | return tfHandle1; | |
| 77 | - break; | ||
| 78 | 73 | default: | |
| 79 | 74 | throw new NotImplementedException("Marshal.Copy failed."); | |
| 80 | 75 | } | |
| 81 | - | ||
| 82 | - var dataType = ToTFDataType(nd.dtype); | ||
| 83 | - // shape | ||
| 84 | - var dims = nd.shape.Select(x => (long)x).ToArray(); | ||
| 76 | + | ||
| 85 | 77 | // Free the original buffer and set flag | |
| 86 | 78 | Deallocator deallocator = (IntPtr values, IntPtr len, ref bool closure) => | |
| 87 | 79 | { | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -162,12 +162,12 @@ public string save(Session sess, | |||
| 162 | 162 | ||
| 163 | 163 | if (!_is_empty) | |
| 164 | 164 | { | |
| 165 | - var model_checkpoint_path1 = sess.run(_saver_def.SaveTensorName, new FeedItem[] { | ||
| 165 | + model_checkpoint_path = sess.run(_saver_def.SaveTensorName, new FeedItem[] { | ||
| 166 | 166 | new FeedItem(_saver_def.FilenameTensorName, checkpoint_file) | |
| 167 | 167 | }); | |
| 168 | 168 | } | |
| 169 | 169 | ||
| 170 | - throw new NotImplementedException(""); | ||
| 170 | + throw new NotImplementedException("Saver.save"); | ||
| 171 | 171 | ||
| 172 | 172 | return model_checkpoint_path; | |
| 173 | 173 | } | |
| Back | FazBrowse Home | New Git URL |
0 commit comments