| FazBrowse GitHub Viewer | Trending | | Home |
| Tools: [Download Repo ZIP] [Original HTTPS Page] |
1 parent f33b203 commit 170f3a7
3 files changed
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -53,6 +53,12 @@ public Saver Saver(VariableV1[] var_list = null, int max_to_keep = 5) | |||
| 53 | 53 | public string write_graph(Graph graph, string logdir, string name, bool as_text = true) | |
| 54 | 54 | => graph_io.write_graph(graph, logdir, name, as_text); | |
| 55 | 55 | ||
| 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 | + | ||
| 56 | 62 | public Saver import_meta_graph(string meta_graph_or_file, | |
| 57 | 63 | bool clear_devices = false, | |
| 58 | 64 | string import_scope = "") => saver._import_meta_graph_with_return_elements(meta_graph_or_file, | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -5,7 +5,7 @@ | |||
| 5 | 5 | <AssemblyName>TensorFlow.NET</AssemblyName> | |
| 6 | 6 | <RootNamespace>Tensorflow</RootNamespace> | |
| 7 | 7 | <TargetTensorFlow>1.14.1</TargetTensorFlow> | |
| 8 | - <Version>0.15.0</Version> | ||
| 8 | + <Version>0.14.1</Version> | ||
| 9 | 9 | <Authors>Haiping Chen, Meinrad Recheis, Eli Belash</Authors> | |
| 10 | 10 | <Company>SciSharp STACK</Company> | |
| 11 | 11 | <GeneratePackageOnBuild>true</GeneratePackageOnBuild> | |
@@ -18,11 +18,12 @@ | |||
| 18 | 18 | <Description>Google's TensorFlow full binding in .NET Standard. | |
| 19 | 19 | Building, training and infering deep learning models. | |
| 20 | 20 | https://tensorflownet.readthedocs.io</Description> | |
| 21 | - <AssemblyVersion>0.15.0.0</AssemblyVersion> | ||
| 21 | + <AssemblyVersion>0.14.1.0</AssemblyVersion> | ||
| 22 | 22 | <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> | ||
| 24 | 25 | <LangVersion>7.3</LangVersion> | |
| 25 | - <FileVersion>0.15.0.0</FileVersion> | ||
| 26 | + <FileVersion>0.14.1.0</FileVersion> | ||
| 26 | 27 | <PackageLicenseFile>LICENSE</PackageLicenseFile> | |
| 27 | 28 | <PackageRequireLicenseAcceptance>true</PackageRequireLicenseAcceptance> | |
| 28 | 29 | <SignAssembly>true</SignAssembly> | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -14,9 +14,12 @@ You may obtain a copy of the License at | |||
| 14 | 14 | limitations under the License. | |
| 15 | 15 | ******************************************************************************/ | |
| 16 | 16 | ||
| 17 | + using Google.Protobuf; | ||
| 17 | 18 | using System; | |
| 18 | 19 | using System.Collections.Generic; | |
| 20 | + using System.IO; | ||
| 19 | 21 | using System.Linq; | |
| 22 | + using static Tensorflow.Binding; | ||
| 20 | 23 | ||
| 21 | 24 | namespace Tensorflow | |
| 22 | 25 | { | |
@@ -81,5 +84,35 @@ public static Saver _create_saver_from_imported_meta_graph(MetaGraphDef meta_gra | |||
| 81 | 84 | } | |
| 82 | 85 | } | |
| 83 | 86 | } | |
| 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 | + } | ||
| 84 | 117 | } | |
| 85 | 118 | } | |
| Back | FazBrowse Home | New Git URL |
0 commit comments