| FazBrowse GitHub Viewer | Trending | | Home |
| Tools: [Download Repo ZIP] [Original HTTPS Page] |
1 parent de082e1 commit 3d7ff13
12 files changed
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -0,0 +1,12 @@ | |||
| 1 | + using System; | ||
| 2 | + using System.Collections.Generic; | ||
| 3 | + using System.Text; | ||
| 4 | + | ||
| 5 | + namespace Tensorflow | ||
| 6 | + { | ||
| 7 | + public static partial class tf | ||
| 8 | + { | ||
| 9 | + public static object get_collection(string key, string scope = "") => get_default_graph() | ||
| 10 | + .get_collection(key, scope: scope); | ||
| 11 | + } | ||
| 12 | + } | ||
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -0,0 +1,10 @@ | |||
| 1 | + using System; | ||
| 2 | + using System.Collections.Generic; | ||
| 3 | + using System.Text; | ||
| 4 | + | ||
| 5 | + namespace Tensorflow.Operations.Losses | ||
| 6 | + { | ||
| 7 | + class losses_impl | ||
| 8 | + { | ||
| 9 | + } | ||
| 10 | + } | ||
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -82,9 +82,9 @@ public static Tensor shape(Tensor input, string name = "", TF_DataType out_type | |||
| 82 | 82 | return shape_internal(input, name, optimize: true, out_type: out_type); | |
| 83 | 83 | } | |
| 84 | 84 | ||
| 85 | - public static Tensor size(Tensor input, string name = "", TF_DataType out_type = TF_DataType.TF_INT32) | ||
| 85 | + public static Tensor size(Tensor input, string name = "", bool optimize = true, TF_DataType out_type = TF_DataType.TF_INT32) | ||
| 86 | 86 | { | |
| 87 | - return size_internal(input, name, optimize: true, out_type: out_type); | ||
| 87 | + return size_internal(input, name, optimize: optimize, out_type: out_type); | ||
| 88 | 88 | } | |
| 89 | 89 | ||
| 90 | 90 | private static Tensor shape_internal(Tensor input, string name = "", bool optimize = true, TF_DataType out_type = TF_DataType.TF_INT32) | |
@@ -132,6 +132,7 @@ private static Tensor size_internal(Tensor input, string name = "", bool optimiz | |||
| 132 | 132 | else | |
| 133 | 133 | { | |
| 134 | 134 | // result = gen_array_ops.shape(); | |
| 135 | + throw new NotImplementedException("array_ops.size_internal"); | ||
| 135 | 136 | } | |
| 136 | 137 | ||
| 137 | 138 | return null; | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -46,28 +46,36 @@ private NDArray _run(object fetches, FeedItem[] feed_dict = null) | |||
| 46 | 46 | var feed_dict_tensor = new Dictionary<object, object>(); | |
| 47 | 47 | var feed_map = new Dictionary<object, object>(); | |
| 48 | 48 | ||
| 49 | + Func<FeedItem, IEnumerable<(object, object)>> feed_fn = (item) => | ||
| 50 | + { | ||
| 51 | + return new (object, object)[] { (item.Key, item.Value) }; | ||
| 52 | + }; | ||
| 53 | + | ||
| 49 | 54 | // Validate and process feed_dict. | |
| 50 | 55 | if (feed_dict != null) | |
| 51 | 56 | { | |
| 52 | - foreach(var subfeed in feed_dict) | ||
| 57 | + foreach (var feed in feed_dict) | ||
| 53 | 58 | { | |
| 54 | - var subfeed_t = _graph.as_graph_element(subfeed.Key, allow_tensor: true, allow_operation: false); | ||
| 55 | - var subfeed_dtype = subfeed_t.dtype.as_numpy_datatype(); | ||
| 56 | - switch(subfeed.Value) | ||
| 59 | + foreach (var (subfeed, subfeed_val) in feed_fn(feed)) | ||
| 57 | 60 | { | |
| 58 | - case float floatVal: | ||
| 59 | - feed_dict_tensor[subfeed_t] = (NDArray)floatVal; | ||
| 60 | - break; | ||
| 61 | - case int intVal: | ||
| 62 | - feed_dict_tensor[subfeed_t] = (NDArray)intVal; | ||
| 63 | - break; | ||
| 64 | - case string str: | ||
| 65 | - feed_dict_tensor[subfeed_t] = (NDArray)str; | ||
| 66 | - break; | ||
| 67 | - default: | ||
| 68 | - throw new NotImplementedException("_run subfeed"); | ||
| 61 | + var subfeed_t = _graph.as_graph_element(subfeed, allow_tensor: true, allow_operation: false); | ||
| 62 | + var subfeed_dtype = subfeed_t.dtype.as_numpy_datatype(); | ||
| 63 | + switch (subfeed_val) | ||
| 64 | + { | ||
| 65 | + case float floatVal: | ||
| 66 | + feed_dict_tensor[subfeed_t] = (NDArray)floatVal; | ||
| 67 | + break; | ||
| 68 | + case int intVal: | ||
| 69 | + feed_dict_tensor[subfeed_t] = (NDArray)intVal; | ||
| 70 | + break; | ||
| 71 | + case string str: | ||
| 72 | + feed_dict_tensor[subfeed_t] = (NDArray)str; | ||
| 73 | + break; | ||
| 74 | + default: | ||
| 75 | + throw new NotImplementedException("_run subfeed"); | ||
| 76 | + } | ||
| 77 | + feed_map[subfeed_t.name] = (subfeed_t, subfeed_val); | ||
| 69 | 78 | } | |
| 70 | - feed_map[subfeed_t.name] = (subfeed_t, subfeed.Value); | ||
| 71 | 79 | } | |
| 72 | 80 | } | |
| 73 | 81 | ||
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -24,7 +24,7 @@ public static Tensor constant(object value, TF_DataType dtype = TF_DataType.DtIn | |||
| 24 | 24 | return _constant_impl(value, dtype, shape, name, verify_shape: false, allow_broadcast: true); | |
| 25 | 25 | } | |
| 26 | 26 | ||
| 27 | - private static Tensor _constant_impl(object value, TF_DataType dtype, int[] shape, string name, bool verify_shape, bool allow_broadcast) | ||
| 27 | + public static Tensor _constant_impl(object value, TF_DataType dtype, int[] shape, string name, bool verify_shape, bool allow_broadcast) | ||
| 28 | 28 | { | |
| 29 | 29 | if (tf.context.executing_eagerly()) | |
| 30 | 30 | { | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -7,8 +7,26 @@ namespace Tensorflow | |||
| 7 | 7 | { | |
| 8 | 8 | public static partial class tf | |
| 9 | 9 | { | |
| 10 | - public static Tensor constant(NDArray nd, string name = "Const") => constant_op.constant(nd, name: name); | ||
| 10 | + // public static Tensor constant(NDArray nd, string name = "Const") => constant_op.constant(nd, name: name); | ||
| 11 | + | ||
| 12 | + public static Tensor constant(object value, | ||
| 13 | + TF_DataType dtype = TF_DataType.DtInvalid, | ||
| 14 | + int[] shape = null, | ||
| 15 | + string name = "Const", | ||
| 16 | + bool verify_shape = false) => constant_op._constant_impl(value, | ||
| 17 | + dtype, | ||
| 18 | + shape, | ||
| 19 | + name, | ||
| 20 | + verify_shape: verify_shape, | ||
| 21 | + allow_broadcast: false); | ||
| 11 | 22 | ||
| 12 | 23 | public static Tensor zeros(Shape shape, TF_DataType dtype = TF_DataType.TF_FLOAT, string name = "") => array_ops.zeros(shape, dtype, name); | |
| 24 | + | ||
| 25 | + public static Tensor size(Tensor input, | ||
| 26 | + string name = "", | ||
| 27 | + TF_DataType out_type = TF_DataType.TF_INT32) => array_ops.size(input, | ||
| 28 | + name, | ||
| 29 | + optimize: true, | ||
| 30 | + out_type: out_type); | ||
| 13 | 31 | } | |
| 14 | 32 | } | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -55,6 +55,7 @@ public Saver(RefVariable[] var_list = null, | |||
| 55 | 55 | _keep_checkpoint_every_n_hours = keep_checkpoint_every_n_hours; | |
| 56 | 56 | _name = name; | |
| 57 | 57 | _restore_sequentially = restore_sequentially; | |
| 58 | + _saver_def = saver_def; | ||
| 58 | 59 | _builder = builder; | |
| 59 | 60 | _is_built = false; | |
| 60 | 61 | _allow_empty = allow_empty; | |
@@ -122,7 +123,7 @@ private void _build(string checkpoint_path, bool build_save, bool build_restore) | |||
| 122 | 123 | } | |
| 123 | 124 | else if (_saver_def != null && !string.IsNullOrEmpty(_name)) | |
| 124 | 125 | { | |
| 125 | - throw new NotImplementedException(""); | ||
| 126 | + throw new NotImplementedException("Saver._build"); | ||
| 126 | 127 | } | |
| 127 | 128 | ||
| 128 | 129 | _check_saver_def(); | |
@@ -200,6 +201,38 @@ public string save(Session sess, | |||
| 200 | 201 | return saver._import_meta_graph_with_return_elements(meta_graph_or_file, clear_devices, import_scope); | |
| 201 | 202 | } | |
| 202 | 203 | ||
| 204 | + /// <summary> | ||
| 205 | + /// Restores previously saved variables. | ||
| 206 | + /// | ||
| 207 | + /// This method runs the ops added by the constructor for restoring variables. | ||
| 208 | + /// It requires a session in which the graph was launched. The variables to | ||
| 209 | + /// restore do not have to have been initialized, as restoring is itself a way | ||
| 210 | + /// to initialize variables. | ||
| 211 | + /// </summary> | ||
| 212 | + /// <param name="sess">A `Session` to use to restore the parameters. None in eager mode.</param> | ||
| 213 | + /// <param name="save_path">Path where parameters were previously saved.</param> | ||
| 214 | + public void restore(Session sess, string save_path) | ||
| 215 | + { | ||
| 216 | + if (_is_empty) | ||
| 217 | + return; | ||
| 218 | + | ||
| 219 | + if (string.IsNullOrEmpty(save_path)) | ||
| 220 | + throw new ValueError("Can't load save_path when it is None."); | ||
| 221 | + | ||
| 222 | + if (!checkpoint_management.checkpoint_exists(save_path)) | ||
| 223 | + throw new ValueError($"The passed save_path is not a valid checkpoint: {save_path}"); | ||
| 224 | + | ||
| 225 | + Console.WriteLine($"Restoring parameters from {save_path}"); | ||
| 226 | + | ||
| 227 | + if (tf.context.executing_eagerly()) | ||
| 228 | + ; | ||
| 229 | + else | ||
| 230 | + sess.run(_saver_def.RestoreOpName, new FeedItem[] | ||
| 231 | + { | ||
| 232 | + new FeedItem(_saver_def.FilenameTensorName, save_path) | ||
| 233 | + }); | ||
| 234 | + } | ||
| 235 | + | ||
| 203 | 236 | /// <summary> | |
| 204 | 237 | /// Writes `MetaGraphDef` to save_path/filename. | |
| 205 | 238 | /// </summary> | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -3,6 +3,7 @@ | |||
| 3 | 3 | using System.IO; | |
| 4 | 4 | using System.Linq; | |
| 5 | 5 | using System.Text; | |
| 6 | + using static Tensorflow.SaverDef.Types; | ||
| 6 | 7 | ||
| 7 | 8 | namespace Tensorflow | |
| 8 | 9 | { | |
@@ -105,5 +106,23 @@ public static string meta_graph_filename(string checkpoint_filename, string meta | |||
| 105 | 106 | string suffixed_filename = basename + "." + meta_graph_suffix; | |
| 106 | 107 | return suffixed_filename; | |
| 107 | 108 | } | |
| 109 | + | ||
| 110 | + public static bool checkpoint_exists(string checkpoint_prefix) | ||
| 111 | + { | ||
| 112 | + string pathname = _prefix_to_checkpoint_path(checkpoint_prefix, CheckpointFormatVersion.V2); | ||
| 113 | + if (File.Exists(pathname)) | ||
| 114 | + return true; | ||
| 115 | + else if (File.Exists(checkpoint_prefix)) | ||
| 116 | + return true; | ||
| 117 | + else | ||
| 118 | + return false; | ||
| 119 | + } | ||
| 120 | + | ||
| 121 | + private static string _prefix_to_checkpoint_path(string prefix, CheckpointFormatVersion format_version) | ||
| 122 | + { | ||
| 123 | + if (format_version == CheckpointFormatVersion.V2) | ||
| 124 | + return prefix + ".index"; | ||
| 125 | + return prefix; | ||
| 126 | + } | ||
| 108 | 127 | } | |
| 109 | 128 | } | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -1,5 +1,6 @@ | |||
| 1 | 1 | using System; | |
| 2 | 2 | using System.Collections.Generic; | |
| 3 | + using System.Linq; | ||
| 3 | 4 | using System.Text; | |
| 4 | 5 | ||
| 5 | 6 | namespace Tensorflow | |
@@ -13,25 +14,43 @@ public static (Saver, object) _import_meta_graph_with_return_elements(string met | |||
| 13 | 14 | { | |
| 14 | 15 | var meta_graph_def = meta_graph.read_meta_graph_file(meta_graph_or_file); | |
| 15 | 16 | ||
| 16 | - var imported_vars = meta_graph.import_scoped_meta_graph_with_return_elements( | ||
| 17 | + var meta = meta_graph.import_scoped_meta_graph_with_return_elements( | ||
| 17 | 18 | meta_graph_def, | |
| 18 | 19 | clear_devices: clear_devices, | |
| 19 | 20 | import_scope: import_scope, | |
| 20 | 21 | return_elements: return_elements); | |
| 21 | 22 | ||
| 23 | + var (imported_vars, imported_return_elements) = meta; | ||
| 24 | + | ||
| 22 | 25 | var saver = _create_saver_from_imported_meta_graph( | |
| 23 | 26 | meta_graph_def, import_scope, imported_vars); | |
| 24 | 27 | ||
| 25 | 28 | return (saver, null); | |
| 26 | 29 | } | |
| 27 | 30 | ||
| 31 | + /// <summary> | ||
| 32 | + /// Return a saver for restoring variable values to an imported MetaGraph. | ||
| 33 | + /// </summary> | ||
| 34 | + /// <param name="meta_graph_def"></param> | ||
| 35 | + /// <param name="import_scope"></param> | ||
| 36 | + /// <param name="imported_vars"></param> | ||
| 37 | + /// <returns></returns> | ||
| 28 | 38 | public static Saver _create_saver_from_imported_meta_graph(MetaGraphDef meta_graph_def, | |
| 29 | 39 | string import_scope, | |
| 30 | - (Dictionary<string, RefVariable>, ITensorOrOperation[]) imported_vars) | ||
| 40 | + Dictionary<string, RefVariable> imported_vars) | ||
| 31 | 41 | { | |
| 32 | 42 | if(meta_graph_def.SaverDef != null) | |
| 33 | 43 | { | |
| 34 | - throw new NotImplementedException("_create_saver_from_imported_meta_graph"); | ||
| 44 | + // Infer the scope that is prepended by `import_scoped_meta_graph`. | ||
| 45 | + string scope = import_scope; | ||
| 46 | + var var_names = imported_vars.Keys.ToArray(); | ||
| 47 | + if(var_names.Length > 0) | ||
| 48 | + { | ||
| 49 | + var sample_key = var_names[0]; | ||
| 50 | + var sample_var = imported_vars[sample_key]; | ||
| 51 | + scope = string.Join("", sample_var.name.Skip(sample_key.Length)); | ||
| 52 | + } | ||
| 53 | + return new Saver(saver_def: meta_graph_def.SaverDef, name: scope); | ||
| 35 | 54 | } | |
| 36 | 55 | else | |
| 37 | 56 | { | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -0,0 +1,29 @@ | |||
| 1 | + using System; | ||
| 2 | + using System.Collections.Generic; | ||
| 3 | + using System.Text; | ||
| 4 | + using Tensorflow; | ||
| 5 | + | ||
| 6 | + namespace TensorFlowNET.Examples | ||
| 7 | + { | ||
| 8 | + public class MetaGraph : Python, IExample | ||
| 9 | + { | ||
| 10 | + public void Run() | ||
| 11 | + { | ||
| 12 | + ImportMetaGraph("my-save-dir/"); | ||
| 13 | + } | ||
| 14 | + | ||
| 15 | + private void ImportMetaGraph(string dir) | ||
| 16 | + { | ||
| 17 | + with<Session>(tf.Session(), sess => | ||
| 18 | + { | ||
| 19 | + var new_saver = tf.train.import_meta_graph(dir + "my-model-10000.meta"); | ||
| 20 | + new_saver.restore(sess, dir + "my-model-10000"); | ||
| 21 | + var labels = tf.constant(0, dtype: tf.int32, shape: new int[] { 100 }, name: "labels"); | ||
| 22 | + var batch_size = tf.size(labels); | ||
| 23 | + var logits = (tf.get_collection("logits") as List<ITensorOrOperation>)[0]; | ||
| 24 | + var loss = tf.losses.sparse_softmax_cross_entropy(labels = labels, | ||
| 25 | + logits = logits); | ||
| 26 | + }); | ||
| 27 | + } | ||
| 28 | + } | ||
| 29 | + } | ||
| Back | FazBrowse Home | New Git URL |
0 commit comments