| FazBrowse GitHub Viewer | Trending | | Home |
| Tools: [Download Repo ZIP] [Original HTTPS Page] |
1 parent 09600a8 commit f7e61b0
7 files changed
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -62,7 +62,7 @@ public static ITensorOrOperation[] import_graph_def(GraphDef graph_def, | |||
| 62 | 62 | { | |
| 63 | 63 | _PopulateTFImportGraphDefOptions(scoped_options, prefix, input_map, return_elements); | |
| 64 | 64 | // need to create a class ImportGraphDefWithResults with IDisposal | |
| 65 | - results = c_api.TF_GraphImportGraphDefWithResults(graph, buffer.Handle, scoped_options.Handle, status.Handle); | ||
| 65 | + results = new TF_ImportGraphDefResults(c_api.TF_GraphImportGraphDefWithResults(graph, buffer.Handle, scoped_options.Handle, status.Handle)); | ||
| 66 | 66 | status.Check(true); | |
| 67 | 67 | } | |
| 68 | 68 | ||
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -33,7 +33,7 @@ public OperationDescription NewOperation(string opType, string opName) | |||
| 33 | 33 | return c_api.TF_NewOperation(_handle, opType, opName); | |
| 34 | 34 | } | |
| 35 | 35 | ||
| 36 | - public Operation[] ReturnOperations(IntPtr results) | ||
| 36 | + public Operation[] ReturnOperations(SafeImportGraphDefResultsHandle results) | ||
| 37 | 37 | { | |
| 38 | 38 | TF_Operation return_oper_handle = new TF_Operation(); | |
| 39 | 39 | int num_return_opers = 0; | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -413,7 +413,7 @@ public string unique_name(string name, bool mark_as_used = true) | |||
| 413 | 413 | return name; | |
| 414 | 414 | } | |
| 415 | 415 | ||
| 416 | - public TF_Output[] ReturnOutputs(IntPtr results) | ||
| 416 | + public TF_Output[] ReturnOutputs(SafeImportGraphDefResultsHandle results) | ||
| 417 | 417 | { | |
| 418 | 418 | IntPtr return_output_handle = IntPtr.Zero; | |
| 419 | 419 | int num_return_outputs = 0; | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -0,0 +1,40 @@ | |||
| 1 | + /***************************************************************************** | ||
| 2 | + Copyright 2018 The TensorFlow.NET Authors. All Rights Reserved. | ||
| 3 | + | ||
| 4 | + Licensed under the Apache License, Version 2.0 (the "License"); | ||
| 5 | + you may not use this file except in compliance with the License. | ||
| 6 | + You may obtain a copy of the License at | ||
| 7 | + | ||
| 8 | + http://www.apache.org/licenses/LICENSE-2.0 | ||
| 9 | + | ||
| 10 | + Unless required by applicable law or agreed to in writing, software | ||
| 11 | + distributed under the License is distributed on an "AS IS" BASIS, | ||
| 12 | + WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. | ||
| 13 | + See the License for the specific language governing permissions and | ||
| 14 | + limitations under the License. | ||
| 15 | + ******************************************************************************/ | ||
| 16 | + | ||
| 17 | + using System; | ||
| 18 | + using Tensorflow.Util; | ||
| 19 | + | ||
| 20 | + namespace Tensorflow | ||
| 21 | + { | ||
| 22 | + public sealed class SafeImportGraphDefResultsHandle : SafeTensorflowHandle | ||
| 23 | + { | ||
| 24 | + private SafeImportGraphDefResultsHandle() | ||
| 25 | + { | ||
| 26 | + } | ||
| 27 | + | ||
| 28 | + public SafeImportGraphDefResultsHandle(IntPtr handle) | ||
| 29 | + : base(handle) | ||
| 30 | + { | ||
| 31 | + } | ||
| 32 | + | ||
| 33 | + protected override bool ReleaseHandle() | ||
| 34 | + { | ||
| 35 | + c_api.TF_DeleteImportGraphDefResults(handle); | ||
| 36 | + SetHandle(IntPtr.Zero); | ||
| 37 | + return true; | ||
| 38 | + } | ||
| 39 | + } | ||
| 40 | + } | ||
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -1,18 +1,35 @@ | |||
| 1 | - using System; | ||
| 2 | - using System.Runtime.InteropServices; | ||
| 1 | + /***************************************************************************** | ||
| 2 | + Copyright 2018 The TensorFlow.NET Authors. All Rights Reserved. | ||
| 3 | + | ||
| 4 | + Licensed under the Apache License, Version 2.0 (the "License"); | ||
| 5 | + you may not use this file except in compliance with the License. | ||
| 6 | + You may obtain a copy of the License at | ||
| 7 | + | ||
| 8 | + http://www.apache.org/licenses/LICENSE-2.0 | ||
| 9 | + | ||
| 10 | + Unless required by applicable law or agreed to in writing, software | ||
| 11 | + distributed under the License is distributed on an "AS IS" BASIS, | ||
| 12 | + WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. | ||
| 13 | + See the License for the specific language governing permissions and | ||
| 14 | + limitations under the License. | ||
| 15 | + ******************************************************************************/ | ||
| 16 | + | ||
| 17 | + using System; | ||
| 3 | 18 | ||
| 4 | 19 | namespace Tensorflow | |
| 5 | 20 | { | |
| 6 | - public class TF_ImportGraphDefResults : DisposableObject | ||
| 21 | + public sealed class TF_ImportGraphDefResults : IDisposable | ||
| 7 | 22 | { | |
| 8 | 23 | /*public IntPtr return_nodes; | |
| 9 | 24 | public IntPtr missing_unused_key_names; | |
| 10 | 25 | public IntPtr missing_unused_key_indexes; | |
| 11 | 26 | public IntPtr missing_unused_key_names_data;*/ | |
| 12 | 27 | ||
| 13 | - public TF_ImportGraphDefResults(IntPtr handle) | ||
| 28 | + private SafeImportGraphDefResultsHandle Handle { get; } | ||
| 29 | + | ||
| 30 | + public TF_ImportGraphDefResults(SafeImportGraphDefResultsHandle handle) | ||
| 14 | 31 | { | |
| 15 | - _handle = handle; | ||
| 32 | + Handle = handle; | ||
| 16 | 33 | } | |
| 17 | 34 | ||
| 18 | 35 | public TF_Output[] return_tensors | |
@@ -21,7 +38,7 @@ public TF_Output[] return_tensors | |||
| 21 | 38 | { | |
| 22 | 39 | IntPtr return_output_handle = IntPtr.Zero; | |
| 23 | 40 | int num_outputs = -1; | |
| 24 | - c_api.TF_ImportGraphDefResultsReturnOutputs(_handle, ref num_outputs, ref return_output_handle); | ||
| 41 | + c_api.TF_ImportGraphDefResultsReturnOutputs(Handle, ref num_outputs, ref return_output_handle); | ||
| 25 | 42 | TF_Output[] return_outputs = new TF_Output[num_outputs]; | |
| 26 | 43 | unsafe | |
| 27 | 44 | { | |
@@ -52,13 +69,7 @@ public TF_Operation[] return_opers | |||
| 52 | 69 | } | |
| 53 | 70 | } | |
| 54 | 71 | ||
| 55 | - public static implicit operator TF_ImportGraphDefResults(IntPtr handle) | ||
| 56 | - => new TF_ImportGraphDefResults(handle); | ||
| 57 | - | ||
| 58 | - public static implicit operator IntPtr(TF_ImportGraphDefResults results) | ||
| 59 | - => results._handle; | ||
| 60 | - | ||
| 61 | - protected override void DisposeUnmanagedResources(IntPtr handle) | ||
| 62 | - => c_api.TF_DeleteImportGraphDefResults(handle); | ||
| 72 | + public void Dispose() | ||
| 73 | + => Handle.Dispose(); | ||
| 63 | 74 | } | |
| 64 | 75 | } | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -92,7 +92,7 @@ public partial class c_api | |||
| 92 | 92 | /// <param name="status">TF_Status*</param> | |
| 93 | 93 | /// <returns>TF_ImportGraphDefResults*</returns> | |
| 94 | 94 | [DllImport(TensorFlowLibName)] | |
| 95 | - public static extern IntPtr TF_GraphImportGraphDefWithResults(IntPtr graph, SafeBufferHandle graph_def, SafeImportGraphDefOptionsHandle options, SafeStatusHandle status); | ||
| 95 | + public static extern SafeImportGraphDefResultsHandle TF_GraphImportGraphDefWithResults(IntPtr graph, SafeBufferHandle graph_def, SafeImportGraphDefOptionsHandle options, SafeStatusHandle status); | ||
| 96 | 96 | ||
| 97 | 97 | /// <summary> | |
| 98 | 98 | /// Import the graph serialized in `graph_def` into `graph`. | |
@@ -258,7 +258,7 @@ public partial class c_api | |||
| 258 | 258 | /// <param name="num_opers">int*</param> | |
| 259 | 259 | /// <param name="opers">TF_Operation***</param> | |
| 260 | 260 | [DllImport(TensorFlowLibName)] | |
| 261 | - public static extern void TF_ImportGraphDefResultsReturnOperations(IntPtr results, ref int num_opers, ref TF_Operation opers); | ||
| 261 | + public static extern void TF_ImportGraphDefResultsReturnOperations(SafeImportGraphDefResultsHandle results, ref int num_opers, ref TF_Operation opers); | ||
| 262 | 262 | ||
| 263 | 263 | /// <summary> | |
| 264 | 264 | /// Fetches the return outputs requested via | |
@@ -270,7 +270,7 @@ public partial class c_api | |||
| 270 | 270 | /// <param name="num_outputs">int*</param> | |
| 271 | 271 | /// <param name="outputs">TF_Output**</param> | |
| 272 | 272 | [DllImport(TensorFlowLibName)] | |
| 273 | - public static extern void TF_ImportGraphDefResultsReturnOutputs(IntPtr results, ref int num_outputs, ref IntPtr outputs); | ||
| 273 | + public static extern void TF_ImportGraphDefResultsReturnOutputs(SafeImportGraphDefResultsHandle results, ref int num_outputs, ref IntPtr outputs); | ||
| 274 | 274 | ||
| 275 | 275 | /// <summary> | |
| 276 | 276 | /// This function creates a new TF_Session (which is created on success) using | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -258,44 +258,49 @@ public void ImportGraphDef() | |||
| 258 | 258 | EXPECT_EQ(0, neg.NumControlOutputs); | |
| 259 | 259 | EXPECT_EQ(0, neg.GetControlOutputs().Length); | |
| 260 | 260 | ||
| 261 | - // Import it again, with an input mapping, return outputs, and a return | ||
| 262 | - // operation, into the same graph. | ||
| 263 | - IntPtr results; | ||
| 264 | - using (var opts = c_api.TF_NewImportGraphDefOptions()) | ||
| 261 | + static SafeImportGraphDefResultsHandle ImportGraph(Status s, Graph graph, Buffer graph_def, Operation scalar) | ||
| 265 | 262 | { | |
| 263 | + using var opts = c_api.TF_NewImportGraphDefOptions(); | ||
| 266 | 264 | c_api.TF_ImportGraphDefOptionsSetPrefix(opts, "imported2"); | |
| 267 | 265 | c_api.TF_ImportGraphDefOptionsAddInputMapping(opts, "scalar", 0, new TF_Output(scalar, 0)); | |
| 268 | 266 | c_api.TF_ImportGraphDefOptionsAddReturnOutput(opts, "feed", 0); | |
| 269 | 267 | c_api.TF_ImportGraphDefOptionsAddReturnOutput(opts, "scalar", 0); | |
| 270 | 268 | EXPECT_EQ(2, c_api.TF_ImportGraphDefOptionsNumReturnOutputs(opts)); | |
| 271 | 269 | c_api.TF_ImportGraphDefOptionsAddReturnOperation(opts, "scalar"); | |
| 272 | 270 | EXPECT_EQ(1, c_api.TF_ImportGraphDefOptionsNumReturnOperations(opts)); | |
| 273 | - results = c_api.TF_GraphImportGraphDefWithResults(graph, graph_def.Handle, opts, s.Handle); | ||
| 271 | + var results = c_api.TF_GraphImportGraphDefWithResults(graph, graph_def.Handle, opts, s.Handle); | ||
| 274 | 272 | EXPECT_EQ(TF_Code.TF_OK, s.Code); | |
| 275 | - } | ||
| 276 | - | ||
| 277 | - Operation scalar2 = graph.OperationByName("imported2/scalar"); | ||
| 278 | - Operation feed2 = graph.OperationByName("imported2/feed"); | ||
| 279 | - Operation neg2 = graph.OperationByName("imported2/neg"); | ||
| 280 | - | ||
| 281 | - // Check input mapping | ||
| 282 | - neg_input = neg.Input(0); | ||
| 283 | - EXPECT_EQ(scalar, neg_input.oper); | ||
| 284 | - EXPECT_EQ(0, neg_input.index); | ||
| 285 | 273 | ||
| 286 | - // Check return outputs | ||
| 287 | - var return_outputs = graph.ReturnOutputs(results); | ||
| 288 | - ASSERT_EQ(2, return_outputs.Length); | ||
| 289 | - EXPECT_EQ(feed2, return_outputs[0].oper); | ||
| 290 | - EXPECT_EQ(0, return_outputs[0].index); | ||
| 291 | - EXPECT_EQ(scalar, return_outputs[1].oper); // remapped | ||
| 292 | - EXPECT_EQ(0, return_outputs[1].index); | ||
| 274 | + return results; | ||
| 275 | + } | ||
| 293 | 276 | ||
| 294 | - // Check return operation | ||
| 295 | - var return_opers = graph.ReturnOperations(results); | ||
| 296 | - ASSERT_EQ(1, return_opers.Length); | ||
| 297 | - EXPECT_EQ(scalar2, return_opers[0]); // not remapped | ||
| 298 | - c_api.TF_DeleteImportGraphDefResults(results); | ||
| 277 | + // Import it again, with an input mapping, return outputs, and a return | ||
| 278 | + // operation, into the same graph. | ||
| 279 | + Operation feed2; | ||
| 280 | + using (SafeImportGraphDefResultsHandle results = ImportGraph(s, graph, graph_def, scalar)) | ||
| 281 | + { | ||
| 282 | + Operation scalar2 = graph.OperationByName("imported2/scalar"); | ||
| 283 | + feed2 = graph.OperationByName("imported2/feed"); | ||
| 284 | + Operation neg2 = graph.OperationByName("imported2/neg"); | ||
| 285 | + | ||
| 286 | + // Check input mapping | ||
| 287 | + neg_input = neg.Input(0); | ||
| 288 | + EXPECT_EQ(scalar, neg_input.oper); | ||
| 289 | + EXPECT_EQ(0, neg_input.index); | ||
| 290 | + | ||
| 291 | + // Check return outputs | ||
| 292 | + var return_outputs = graph.ReturnOutputs(results); | ||
| 293 | + ASSERT_EQ(2, return_outputs.Length); | ||
| 294 | + EXPECT_EQ(feed2, return_outputs[0].oper); | ||
| 295 | + EXPECT_EQ(0, return_outputs[0].index); | ||
| 296 | + EXPECT_EQ(scalar, return_outputs[1].oper); // remapped | ||
| 297 | + EXPECT_EQ(0, return_outputs[1].index); | ||
| 298 | + | ||
| 299 | + // Check return operation | ||
| 300 | + var return_opers = graph.ReturnOperations(results); | ||
| 301 | + ASSERT_EQ(1, return_opers.Length); | ||
| 302 | + EXPECT_EQ(scalar2, return_opers[0]); // not remapped | ||
| 303 | + } | ||
| 299 | 304 | ||
| 300 | 305 | // Import again, with control dependencies, into the same graph. | |
| 301 | 306 | using (var opts = c_api.TF_NewImportGraphDefOptions()) | |
| Back | FazBrowse Home | New Git URL |
0 commit comments