| FazBrowse GitHub Viewer | Trending | | Home |
| Tools: [Download Repo ZIP] [Original HTTPS Page] |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -287,7 +287,7 @@ public partial class c_api | |||
| 287 | 287 | /// <param name="status">TF_Status*</param> | |
| 288 | 288 | /// <returns></returns> | |
| 289 | 289 | [DllImport(TensorFlowLibName)] | |
| 290 | - public static extern IntPtr TF_LoadSessionFromSavedModel(IntPtr session_options, IntPtr run_options, | ||
| 290 | + public static extern IntPtr TF_LoadSessionFromSavedModel(SafeSessionOptionsHandle session_options, IntPtr run_options, | ||
| 291 | 291 | string export_dir, string[] tags, int tags_len, | |
| 292 | 292 | IntPtr graph, ref TF_Buffer meta_graph_def, SafeStatusHandle status); | |
| 293 | 293 | ||
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -47,7 +47,7 @@ public BaseSession(string target = "", Graph g = null, ConfigProto config = null | |||
| 47 | 47 | lock (Locks.ProcessWide) | |
| 48 | 48 | { | |
| 49 | 49 | status = status ?? new Status(); | |
| 50 | - _handle = c_api.TF_NewSession(_graph, opts, status.Handle); | ||
| 50 | + _handle = c_api.TF_NewSession(_graph, opts.Handle, status.Handle); | ||
| 51 | 51 | status.Check(true); | |
| 52 | 52 | } | |
| 53 | 53 | } | |
| 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 SafeSessionOptionsHandle : SafeTensorflowHandle | ||
| 23 | + { | ||
| 24 | + public SafeSessionOptionsHandle() | ||
| 25 | + { | ||
| 26 | + } | ||
| 27 | + | ||
| 28 | + public SafeSessionOptionsHandle(IntPtr handle) | ||
| 29 | + : base(handle) | ||
| 30 | + { | ||
| 31 | + } | ||
| 32 | + | ||
| 33 | + protected override bool ReleaseHandle() | ||
| 34 | + { | ||
| 35 | + c_api.TF_DeleteSessionOptions(handle); | ||
| 36 | + SetHandle(IntPtr.Zero); | ||
| 37 | + return true; | ||
| 38 | + } | ||
| 39 | + } | ||
| 40 | + } | ||
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -55,7 +55,7 @@ public static Session LoadFromSavedModel(string path) | |||
| 55 | 55 | IntPtr sess; | |
| 56 | 56 | try | |
| 57 | 57 | { | |
| 58 | - sess = c_api.TF_LoadSessionFromSavedModel(opt, | ||
| 58 | + sess = c_api.TF_LoadSessionFromSavedModel(opt.Handle, | ||
| 59 | 59 | IntPtr.Zero, | |
| 60 | 60 | path, | |
| 61 | 61 | tags, | |
@@ -66,7 +66,7 @@ public static Session LoadFromSavedModel(string path) | |||
| 66 | 66 | status.Check(true); | |
| 67 | 67 | } catch (TensorflowException ex) when (ex.Message.Contains("Could not find SavedModel")) | |
| 68 | 68 | { | |
| 69 | - sess = c_api.TF_LoadSessionFromSavedModel(opt, | ||
| 69 | + sess = c_api.TF_LoadSessionFromSavedModel(opt.Handle, | ||
| 70 | 70 | IntPtr.Zero, | |
| 71 | 71 | Path.GetFullPath(path), | |
| 72 | 72 | tags, | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -16,44 +16,36 @@ limitations under the License. | |||
| 16 | 16 | ||
| 17 | 17 | using Google.Protobuf; | |
| 18 | 18 | using System; | |
| 19 | - using System.Runtime.InteropServices; | ||
| 20 | 19 | ||
| 21 | 20 | namespace Tensorflow | |
| 22 | 21 | { | |
| 23 | - internal class SessionOptions : DisposableObject | ||
| 22 | + internal sealed class SessionOptions : IDisposable | ||
| 24 | 23 | { | |
| 24 | + public SafeSessionOptionsHandle Handle { get; } | ||
| 25 | + | ||
| 25 | 26 | public SessionOptions(string target = "", ConfigProto config = null) | |
| 26 | 27 | { | |
| 27 | - _handle = c_api.TF_NewSessionOptions(); | ||
| 28 | - c_api.TF_SetTarget(_handle, target); | ||
| 28 | + Handle = c_api.TF_NewSessionOptions(); | ||
| 29 | + c_api.TF_SetTarget(Handle, target); | ||
| 29 | 30 | if (config != null) | |
| 30 | 31 | SetConfig(config); | |
| 31 | 32 | } | |
| 32 | 33 | ||
| 33 | - public SessionOptions(IntPtr handle) | ||
| 34 | - { | ||
| 35 | - _handle = handle; | ||
| 36 | - } | ||
| 37 | - | ||
| 38 | - protected override void DisposeUnmanagedResources(IntPtr handle) | ||
| 39 | - => c_api.TF_DeleteSessionOptions(handle); | ||
| 34 | + public void Dispose() | ||
| 35 | + => Handle.Dispose(); | ||
| 40 | 36 | ||
| 41 | - private void SetConfig(ConfigProto config) | ||
| 37 | + private unsafe void SetConfig(ConfigProto config) | ||
| 42 | 38 | { | |
| 43 | 39 | var bytes = config.ToByteArray(); | |
| 44 | - var proto = Marshal.AllocHGlobal(bytes.Length); | ||
| 45 | - Marshal.Copy(bytes, 0, proto, bytes.Length); | ||
| 46 | 40 | ||
| 47 | - using (var status = new Status()) | ||
| 41 | + fixed (byte* proto2 = bytes) | ||
| 48 | 42 | { | |
| 49 | - c_api.TF_SetConfig(_handle, proto, (ulong)bytes.Length, status.Handle); | ||
| 50 | - status.Check(false); | ||
| 43 | + using (var status = new Status()) | ||
| 44 | + { | ||
| 45 | + c_api.TF_SetConfig(Handle, (IntPtr)proto2, (ulong)bytes.Length, status.Handle); | ||
| 46 | + status.Check(false); | ||
| 47 | + } | ||
| 51 | 48 | } | |
| 52 | - | ||
| 53 | - Marshal.FreeHGlobal(proto); | ||
| 54 | 49 | } | |
| 55 | - | ||
| 56 | - public static implicit operator IntPtr(SessionOptions opts) => opts._handle; | ||
| 57 | - public static implicit operator SessionOptions(IntPtr handle) => new SessionOptions(handle); | ||
| 58 | 50 | } | |
| 59 | 51 | } | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -50,14 +50,14 @@ public partial class c_api | |||
| 50 | 50 | /// <param name="status">TF_Status*</param> | |
| 51 | 51 | /// <returns>TF_Session*</returns> | |
| 52 | 52 | [DllImport(TensorFlowLibName)] | |
| 53 | - public static extern IntPtr TF_NewSession(IntPtr graph, IntPtr opts, SafeStatusHandle status); | ||
| 53 | + public static extern IntPtr TF_NewSession(IntPtr graph, SafeSessionOptionsHandle opts, SafeStatusHandle status); | ||
| 54 | 54 | ||
| 55 | 55 | /// <summary> | |
| 56 | 56 | /// Return a new options object. | |
| 57 | 57 | /// </summary> | |
| 58 | 58 | /// <returns>TF_SessionOptions*</returns> | |
| 59 | 59 | [DllImport(TensorFlowLibName)] | |
| 60 | - public static extern unsafe IntPtr TF_NewSessionOptions(); | ||
| 60 | + public static extern SafeSessionOptionsHandle TF_NewSessionOptions(); | ||
| 61 | 61 | ||
| 62 | 62 | /// <summary> | |
| 63 | 63 | /// Run the graph associated with the session starting with the supplied inputs | |
@@ -116,9 +116,9 @@ public static extern unsafe void TF_SessionRun(IntPtr session, TF_Buffer* run_op | |||
| 116 | 116 | /// <param name="proto_len">size_t</param> | |
| 117 | 117 | /// <param name="status">TF_Status*</param> | |
| 118 | 118 | [DllImport(TensorFlowLibName)] | |
| 119 | - public static extern void TF_SetConfig(IntPtr options, IntPtr proto, ulong proto_len, SafeStatusHandle status); | ||
| 119 | + public static extern void TF_SetConfig(SafeSessionOptionsHandle options, IntPtr proto, ulong proto_len, SafeStatusHandle status); | ||
| 120 | 120 | ||
| 121 | 121 | [DllImport(TensorFlowLibName)] | |
| 122 | - public static extern void TF_SetTarget(IntPtr options, string target); | ||
| 122 | + public static extern void TF_SetTarget(SafeSessionOptionsHandle options, string target); | ||
| 123 | 123 | } | |
| 124 | 124 | } | |
| Back | FazBrowse Home | New Git URL |
0 commit comments