FazBrowse GitHub Viewer | Trending |
URL:
| Home
Tools: [Download Repo ZIP]   [Original HTTPS Page]

Implement SafeSessionOptionsHandle as a wrapper for TF_SessionOptions · MSavameri/TensorFlow.NET@bfac6f7 · GitHub

Repository navigation

Commit bfac6f7

Browse files
committed
Implement SafeSessionOptionsHandle as a wrapper for TF_SessionOptions
1 parent 6f74c42 commit bfac6f7

6 files changed

Lines changed: 62 additions & 30 deletions

File tree

‎src/TensorFlowNET.Core/Graphs/c_api.graph.cs‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -287,7 +287,7 @@ public partial class c_api
287287
/// <param name="status">TF_Status*</param>
288288
/// <returns></returns>
289289
[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,
291291
string export_dir, string[] tags, int tags_len,
292292
IntPtr graph, ref TF_Buffer meta_graph_def, SafeStatusHandle status);
293293

‎src/TensorFlowNET.Core/Sessions/BaseSession.cs‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -47,7 +47,7 @@ public BaseSession(string target = "", Graph g = null, ConfigProto config = null
4747
lock (Locks.ProcessWide)
4848
{
4949
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);
5151
status.Check(true);
5252
}
5353
}
Lines changed: 40 additions & 0 deletions
Original file line numberDiff line numberDiff 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+
}

‎src/TensorFlowNET.Core/Sessions/Session.cs‎

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -55,7 +55,7 @@ public static Session LoadFromSavedModel(string path)
5555
IntPtr sess;
5656
try
5757
{
58-
sess = c_api.TF_LoadSessionFromSavedModel(opt,
58+
sess = c_api.TF_LoadSessionFromSavedModel(opt.Handle,
5959
IntPtr.Zero,
6060
path,
6161
tags,
@@ -66,7 +66,7 @@ public static Session LoadFromSavedModel(string path)
6666
status.Check(true);
6767
} catch (TensorflowException ex) when (ex.Message.Contains("Could not find SavedModel"))
6868
{
69-
sess = c_api.TF_LoadSessionFromSavedModel(opt,
69+
sess = c_api.TF_LoadSessionFromSavedModel(opt.Handle,
7070
IntPtr.Zero,
7171
Path.GetFullPath(path),
7272
tags,

‎src/TensorFlowNET.Core/Sessions/SessionOptions.cs‎

Lines changed: 14 additions & 22 deletions
Original file line numberDiff line numberDiff line change
@@ -16,44 +16,36 @@ limitations under the License.
1616

1717
using Google.Protobuf;
1818
using System;
19-
using System.Runtime.InteropServices;
2019

2120
namespace Tensorflow
2221
{
23-
internal class SessionOptions : DisposableObject
22+
internal sealed class SessionOptions : IDisposable
2423
{
24+
public SafeSessionOptionsHandle Handle { get; }
25+
2526
public SessionOptions(string target = "", ConfigProto config = null)
2627
{
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);
2930
if (config != null)
3031
SetConfig(config);
3132
}
3233

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();
4036

41-
private void SetConfig(ConfigProto config)
37+
private unsafe void SetConfig(ConfigProto config)
4238
{
4339
var bytes = config.ToByteArray();
44-
var proto = Marshal.AllocHGlobal(bytes.Length);
45-
Marshal.Copy(bytes, 0, proto, bytes.Length);
4640

47-
using (var status = new Status())
41+
fixed (byte* proto2 = bytes)
4842
{
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+
}
5148
}
52-
53-
Marshal.FreeHGlobal(proto);
5449
}
55-
56-
public static implicit operator IntPtr(SessionOptions opts) => opts._handle;
57-
public static implicit operator SessionOptions(IntPtr handle) => new SessionOptions(handle);
5850
}
5951
}

‎src/TensorFlowNET.Core/Sessions/c_api.session.cs‎

Lines changed: 4 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -50,14 +50,14 @@ public partial class c_api
5050
/// <param name="status">TF_Status*</param>
5151
/// <returns>TF_Session*</returns>
5252
[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);
5454

5555
/// <summary>
5656
/// Return a new options object.
5757
/// </summary>
5858
/// <returns>TF_SessionOptions*</returns>
5959
[DllImport(TensorFlowLibName)]
60-
public static extern unsafe IntPtr TF_NewSessionOptions();
60+
public static extern SafeSessionOptionsHandle TF_NewSessionOptions();
6161

6262
/// <summary>
6363
/// 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
116116
/// <param name="proto_len">size_t</param>
117117
/// <param name="status">TF_Status*</param>
118118
[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);
120120

121121
[DllImport(TensorFlowLibName)]
122-
public static extern void TF_SetTarget(IntPtr options, string target);
122+
public static extern void TF_SetTarget(SafeSessionOptionsHandle options, string target);
123123
}
124124
}

0 commit comments

Comments
 (0)

Back | FazBrowse Home | New Git URL