| FazBrowse GitHub Viewer | Trending | | Home |
| Tools: [Download Repo ZIP] [Original HTTPS Page] |
1 parent 3c25a54 commit 4eccba4
7 files changed
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -47,6 +47,16 @@ public Graph() | |||
| 47 | 47 | _graph_key = $"grap-key-{ops.uid()}/"; | |
| 48 | 48 | } | |
| 49 | 49 | ||
| 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 | + | ||
| 50 | 60 | public ITensorOrOperation as_graph_element(object obj, bool allow_tensor = true, bool allow_operation = true) | |
| 51 | 61 | { | |
| 52 | 62 | return _as_graph_element_locked(obj, allow_tensor, allow_operation); | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -254,6 +254,25 @@ public partial class c_api | |||
| 254 | 254 | [DllImport(TensorFlowLibName)] | |
| 255 | 255 | public static extern void TF_ImportGraphDefResultsReturnOutputs(IntPtr results, ref int num_outputs, ref IntPtr outputs); | |
| 256 | 256 | ||
| 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 | + | ||
| 257 | 276 | [DllImport(TensorFlowLibName)] | |
| 258 | 277 | public static extern IntPtr TF_NewGraph(); | |
| 259 | 278 | ||
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -17,11 +17,23 @@ public Tensor sparse_softmax_cross_entropy(Tensor labels, | |||
| 17 | 17 | (logits, labels, weights)), | |
| 18 | 18 | namescope => | |
| 19 | 19 | { | |
| 20 | - | ||
| 20 | + (labels, logits, weights) = _remove_squeezable_dimensions( | ||
| 21 | + labels, logits, weights, expected_rank_diff: 1); | ||
| 21 | 22 | ||
| 22 | 23 | }); | |
| 23 | 24 | ||
| 24 | 25 | throw new NotImplementedException("sparse_softmax_cross_entropy"); | |
| 25 | 26 | } | |
| 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 | + } | ||
| 26 | 38 | } | |
| 27 | 39 | } | |
| Original file line number | Diff line number | Diff 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 | + } | ||
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -36,6 +36,25 @@ public Session(Graph graph, SessionOptions opts, Status s = null) | |||
| 36 | 36 | Status.Check(true); | |
| 37 | 37 | } | |
| 38 | 38 | ||
| 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 | + | ||
| 39 | 58 | public static implicit operator IntPtr(Session session) => session._handle; | |
| 40 | 59 | public static implicit operator Session(IntPtr handle) => new Session(handle); | |
| 41 | 60 | ||
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -4,7 +4,7 @@ | |||
| 4 | 4 | <TargetFramework>netstandard2.0</TargetFramework> | |
| 5 | 5 | <AssemblyName>TensorFlow.NET</AssemblyName> | |
| 6 | 6 | <RootNamespace>Tensorflow</RootNamespace> | |
| 7 | - <Version>0.2.0</Version> | ||
| 7 | + <Version>0.3.0</Version> | ||
| 8 | 8 | <Authors>Haiping Chen</Authors> | |
| 9 | 9 | <Company>SciSharp STACK</Company> | |
| 10 | 10 | <GeneratePackageOnBuild>true</GeneratePackageOnBuild> | |
@@ -16,13 +16,13 @@ | |||
| 16 | 16 | <PackageTags>TensorFlow, NumSharp, SciSharp, MachineLearning, TensorFlow.NET</PackageTags> | |
| 17 | 17 | <Description>Google's TensorFlow binding in .NET Standard. | |
| 18 | 18 | Docs: https://tensorflownet.readthedocs.io</Description> | |
| 19 | - <AssemblyVersion>0.2.0.0</AssemblyVersion> | ||
| 19 | + <AssemblyVersion>0.3.0.0</AssemblyVersion> | ||
| 20 | 20 | <PackageReleaseNotes>Added a bunch of APIs. | |
| 21 | 21 | Fixed String tensor creation bug. | |
| 22 | 22 | Upgraded to TensorFlow 1.13 RC-1. | |
| 23 | 23 | </PackageReleaseNotes> | |
| 24 | 24 | <LangVersion>7.2</LangVersion> | |
| 25 | - <FileVersion>0.2.0.0</FileVersion> | ||
| 25 | + <FileVersion>0.3.0.0</FileVersion> | ||
| 26 | 26 | </PropertyGroup> | |
| 27 | 27 | ||
| 28 | 28 | <PropertyGroup Condition="'$(Configuration)|$(Platform)'=='Debug|AnyCPU'"> | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -26,6 +26,15 @@ public void ImportGraph() | |||
| 26 | 26 | }); | |
| 27 | 27 | } | |
| 28 | 28 | ||
| 29 | + [TestMethod] | ||
| 30 | + public void ImportSavedModel() | ||
| 31 | + { | ||
| 32 | + with<Session>(Session.LoadFromSavedModel("mobilenet"), sess => | ||
| 33 | + { | ||
| 34 | + | ||
| 35 | + }); | ||
| 36 | + } | ||
| 37 | + | ||
| 29 | 38 | [TestMethod] | |
| 30 | 39 | public void Save1() | |
| 31 | 40 | { | |
| Back | FazBrowse Home | New Git URL |
0 commit comments