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

tf.trainer.load_graph, tf.trainer.freeze_graph · MSavameri/TensorFlow.NET@170f3a7 · GitHub

Repository navigation

Commit 170f3a7

Browse files
committed
tf.trainer.load_graph, tf.trainer.freeze_graph
1 parent f33b203 commit 170f3a7

3 files changed

Lines changed: 44 additions & 4 deletions

File tree

‎src/TensorFlowNET.Core/APIs/tf.train.cs‎

Lines changed: 6 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -53,6 +53,12 @@ public Saver Saver(VariableV1[] var_list = null, int max_to_keep = 5)
5353
public string write_graph(Graph graph, string logdir, string name, bool as_text = true)
5454
=> graph_io.write_graph(graph, logdir, name, as_text);
5555

56+
public Graph load_graph(string freeze_graph_pb)
57+
=> saver.load_graph(freeze_graph_pb);
58+
59+
public string freeze_graph(string checkpoint_dir, string output_pb_name)
60+
=> saver.freeze_graph(checkpoint_dir, output_pb_name);
61+
5662
public Saver import_meta_graph(string meta_graph_or_file,
5763
bool clear_devices = false,
5864
string import_scope = "") => saver._import_meta_graph_with_return_elements(meta_graph_or_file,

‎src/TensorFlowNET.Core/TensorFlow.Binding.csproj‎

Lines changed: 5 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -5,7 +5,7 @@
55
<AssemblyName>TensorFlow.NET</AssemblyName>
66
<RootNamespace>Tensorflow</RootNamespace>
77
<TargetTensorFlow>1.14.1</TargetTensorFlow>
8-
<Version>0.15.0</Version>
8+
<Version>0.14.1</Version>
99
<Authors>Haiping Chen, Meinrad Recheis, Eli Belash</Authors>
1010
<Company>SciSharp STACK</Company>
1111
<GeneratePackageOnBuild>true</GeneratePackageOnBuild>
@@ -18,11 +18,12 @@
1818
<Description>Google's TensorFlow full binding in .NET Standard.
1919
Building, training and infering deep learning models.
2020
https://tensorflownet.readthedocs.io</Description>
21-
<AssemblyVersion>0.15.0.0</AssemblyVersion>
21+
<AssemblyVersion>0.14.1.0</AssemblyVersion>
2222
<PackageReleaseNotes>Changes since v0.14.0:
23-
1: Add TransformGraphWithStringInputs.</PackageReleaseNotes>
23+
1: Add TransformGraphWithStringInputs.
24+
2: tf.trainer.load_graph, tf.trainer.freeze_graph</PackageReleaseNotes>
2425
<LangVersion>7.3</LangVersion>
25-
<FileVersion>0.15.0.0</FileVersion>
26+
<FileVersion>0.14.1.0</FileVersion>
2627
<PackageLicenseFile>LICENSE</PackageLicenseFile>
2728
<PackageRequireLicenseAcceptance>true</PackageRequireLicenseAcceptance>
2829
<SignAssembly>true</SignAssembly>

‎src/TensorFlowNET.Core/Training/Saving/saver.py.cs‎

Lines changed: 33 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -14,9 +14,12 @@ You may obtain a copy of the License at
1414
limitations under the License.
1515
******************************************************************************/
1616

17+
using Google.Protobuf;
1718
using System;
1819
using System.Collections.Generic;
20+
using System.IO;
1921
using System.Linq;
22+
using static Tensorflow.Binding;
2023

2124
namespace Tensorflow
2225
{
@@ -81,5 +84,35 @@ public static Saver _create_saver_from_imported_meta_graph(MetaGraphDef meta_gra
8184
}
8285
}
8386
}
87+
88+
public static string freeze_graph(string checkpoint_dir, string output_pb_name)
89+
{
90+
var checkpoint = checkpoint_management.latest_checkpoint(checkpoint_dir);
91+
if (!File.Exists($"{checkpoint}.meta")) return null;
92+
93+
string output_pb = Path.GetFullPath(Path.Combine(checkpoint_dir, "../", $"{output_pb_name}.pb"));
94+
95+
using (var graph = tf.Graph())
96+
using (var sess = tf.Session(graph))
97+
{
98+
var saver = tf.train.import_meta_graph($"{checkpoint}.meta", clear_devices: true);
99+
saver.restore(sess, checkpoint);
100+
var output_graph_def = tf.graph_util.convert_variables_to_constants(sess,
101+
graph.as_graph_def(),
102+
new string[] { "output/ArgMax" });
103+
Console.WriteLine($"Froze {output_graph_def.Node.Count} nodes.");
104+
File.WriteAllBytes(output_pb, output_graph_def.ToByteArray());
105+
return output_pb;
106+
}
107+
}
108+
109+
public static Graph load_graph(string freeze_graph_pb, string name = "")
110+
{
111+
var bytes = File.ReadAllBytes(freeze_graph_pb);
112+
var graph = tf.Graph().as_default();
113+
importer.import_graph_def(GraphDef.Parser.ParseFrom(bytes),
114+
name: name);
115+
return graph;
116+
}
84117
}
85118
}

0 commit comments

Comments
 (0)

Back | FazBrowse Home | New Git URL