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

add LoadFromSavedModel · attackgithub/TensorFlow.NET@4eccba4 · GitHub

Repository navigation

Commit 4eccba4

Browse files
committed
add LoadFromSavedModel
1 parent 3c25a54 commit 4eccba4

7 files changed

Lines changed: 90 additions & 4 deletions

File tree

‎src/TensorFlowNET.Core/Graphs/Graph.cs‎

Lines changed: 10 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -47,6 +47,16 @@ public Graph()
4747
_graph_key = $"grap-key-{ops.uid()}/";
4848
}
4949

50+
public Graph(IntPtr handle)
51+
{
52+
_handle = handle;
53+
Status = new Status();
54+
_nodes_by_id = new Dictionary<int, ITensorOrOperation>();
55+
_nodes_by_name = new Dictionary<string, ITensorOrOperation>();
56+
_names_in_use = new Dictionary<string, int>();
57+
_graph_key = $"grap-key-{ops.uid()}/";
58+
}
59+
5060
public ITensorOrOperation as_graph_element(object obj, bool allow_tensor = true, bool allow_operation = true)
5161
{
5262
return _as_graph_element_locked(obj, allow_tensor, allow_operation);

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

Lines changed: 19 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -254,6 +254,25 @@ public partial class c_api
254254
[DllImport(TensorFlowLibName)]
255255
public static extern void TF_ImportGraphDefResultsReturnOutputs(IntPtr results, ref int num_outputs, ref IntPtr outputs);
256256

257+
/// <summary>
258+
/// This function creates a new TF_Session (which is created on success) using
259+
/// `session_options`, and then initializes state (restoring tensors and other
260+
/// assets) using `run_options`.
261+
/// </summary>
262+
/// <param name="session_options">const TF_SessionOptions*</param>
263+
/// <param name="run_options">const TF_Buffer*</param>
264+
/// <param name="export_dir">const char*</param>
265+
/// <param name="tags">const char* const*</param>
266+
/// <param name="tags_len">int</param>
267+
/// <param name="graph">TF_Graph*</param>
268+
/// <param name="meta_graph_def">TF_Buffer*</param>
269+
/// <param name="status">TF_Status*</param>
270+
/// <returns></returns>
271+
[DllImport(TensorFlowLibName)]
272+
public static extern IntPtr TF_LoadSessionFromSavedModel(IntPtr session_options, IntPtr run_options,
273+
string export_dir, string[] tags, int tags_len,
274+
IntPtr graph, ref TF_Buffer meta_graph_def, IntPtr status);
275+
257276
[DllImport(TensorFlowLibName)]
258277
public static extern IntPtr TF_NewGraph();
259278

‎src/TensorFlowNET.Core/Operations/Losses/losses_impl.py.cs‎

Lines changed: 13 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -17,11 +17,23 @@ public Tensor sparse_softmax_cross_entropy(Tensor labels,
1717
(logits, labels, weights)),
1818
namescope =>
1919
{
20-
20+
(labels, logits, weights) = _remove_squeezable_dimensions(
21+
labels, logits, weights, expected_rank_diff: 1);
2122

2223
});
2324

2425
throw new NotImplementedException("sparse_softmax_cross_entropy");
2526
}
27+
28+
public (Tensor, Tensor, float) _remove_squeezable_dimensions(Tensor labels,
29+
Tensor predictions,
30+
float weights = 0,
31+
int expected_rank_diff = 0)
32+
{
33+
(labels, predictions, weights) = confusion_matrix.remove_squeezable_dimensions(
34+
labels, predictions, expected_rank_diff: expected_rank_diff);
35+
36+
throw new NotImplementedException("_remove_squeezable_dimensions");
37+
}
2638
}
2739
}
Lines changed: 17 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,17 @@
1+
using System;
2+
using System.Collections.Generic;
3+
using System.Text;
4+
5+
namespace Tensorflow
6+
{
7+
public class confusion_matrix
8+
{
9+
public static (Tensor, Tensor, float) remove_squeezable_dimensions(Tensor labels,
10+
Tensor predictions,
11+
int expected_rank_diff = 0,
12+
string name = "")
13+
{
14+
throw new NotImplementedException("remove_squeezable_dimensions");
15+
}
16+
}
17+
}

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

Lines changed: 19 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -36,6 +36,25 @@ public Session(Graph graph, SessionOptions opts, Status s = null)
3636
Status.Check(true);
3737
}
3838

39+
public static Session LoadFromSavedModel(string path)
40+
{
41+
var graph = c_api.TF_NewGraph();
42+
var status = new Status();
43+
var opt = c_api.TF_NewSessionOptions();
44+
45+
var buffer = new TF_Buffer();
46+
var sess = c_api.TF_LoadSessionFromSavedModel(opt, IntPtr.Zero, path, new string[0], 0, graph, ref buffer, status);
47+
48+
//var bytes = new Buffer(buffer.data).Data;
49+
//var meta_graph = MetaGraphDef.Parser.ParseFrom(bytes);
50+
51+
status.Check();
52+
53+
tf.g = new Graph(graph);
54+
55+
return sess;
56+
}
57+
3958
public static implicit operator IntPtr(Session session) => session._handle;
4059
public static implicit operator Session(IntPtr handle) => new Session(handle);
4160

‎src/TensorFlowNET.Core/TensorFlowNET.Core.csproj‎

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -4,7 +4,7 @@
44
<TargetFramework>netstandard2.0</TargetFramework>
55
<AssemblyName>TensorFlow.NET</AssemblyName>
66
<RootNamespace>Tensorflow</RootNamespace>
7-
<Version>0.2.0</Version>
7+
<Version>0.3.0</Version>
88
<Authors>Haiping Chen</Authors>
99
<Company>SciSharp STACK</Company>
1010
<GeneratePackageOnBuild>true</GeneratePackageOnBuild>
@@ -16,13 +16,13 @@
1616
<PackageTags>TensorFlow, NumSharp, SciSharp, MachineLearning, TensorFlow.NET</PackageTags>
1717
<Description>Google's TensorFlow binding in .NET Standard.
1818
Docs: https://tensorflownet.readthedocs.io</Description>
19-
<AssemblyVersion>0.2.0.0</AssemblyVersion>
19+
<AssemblyVersion>0.3.0.0</AssemblyVersion>
2020
<PackageReleaseNotes>Added a bunch of APIs.
2121
Fixed String tensor creation bug.
2222
Upgraded to TensorFlow 1.13 RC-1.
2323
</PackageReleaseNotes>
2424
<LangVersion>7.2</LangVersion>
25-
<FileVersion>0.2.0.0</FileVersion>
25+
<FileVersion>0.3.0.0</FileVersion>
2626
</PropertyGroup>
2727

2828
<PropertyGroup Condition="'$(Configuration)|$(Platform)'=='Debug|AnyCPU'">

‎test/TensorFlowNET.UnitTest/TrainSaverTest.cs‎

Lines changed: 9 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -26,6 +26,15 @@ public void ImportGraph()
2626
});
2727
}
2828

29+
[TestMethod]
30+
public void ImportSavedModel()
31+
{
32+
with<Session>(Session.LoadFromSavedModel("mobilenet"), sess =>
33+
{
34+
35+
});
36+
}
37+
2938
[TestMethod]
3039
public void Save1()
3140
{

0 commit comments

Comments
 (0)

Back | FazBrowse Home | New Git URL