| FazBrowse GitHub Viewer | Trending | | Home |
| Tools: [Download Repo ZIP] [Original HTTPS Page] |
1 parent 649623b commit d61ccf0
49 files changed
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -17,7 +17,7 @@ public static partial class tf | |||
| 17 | 17 | /// A `Tensor` with the same data as `input`, but its shape has an additional | |
| 18 | 18 | /// dimension of size 1 added. | |
| 19 | 19 | /// </returns> | |
| 20 | - public static Tensor expand_dims(Tensor input, int axis = -1, string name = "", int dim = -1) | ||
| 20 | + public static Tensor expand_dims(Tensor input, int axis = -1, string name = null, int dim = -1) | ||
| 21 | 21 | => array_ops.expand_dims(input, axis, name, dim); | |
| 22 | 22 | } | |
| 23 | 23 | } | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -6,14 +6,14 @@ namespace Tensorflow | |||
| 6 | 6 | { | |
| 7 | 7 | public partial class tf | |
| 8 | 8 | { | |
| 9 | - public static Tensor read_file(string filename, string name = "") => gen_io_ops.read_file(filename, name); | ||
| 9 | + public static Tensor read_file(string filename, string name = null) => gen_io_ops.read_file(filename, name); | ||
| 10 | 10 | ||
| 11 | 11 | public static gen_image_ops image => new gen_image_ops(); | |
| 12 | 12 | ||
| 13 | 13 | public static void import_graph_def(GraphDef graph_def, | |
| 14 | 14 | Dictionary<string, Tensor> input_map = null, | |
| 15 | 15 | string[] return_elements = null, | |
| 16 | - string name = "", | ||
| 16 | + string name = null, | ||
| 17 | 17 | OpList producer_op_list = null) => importer.import_graph_def(graph_def, input_map, return_elements, name, producer_op_list); | |
| 18 | 18 | } | |
| 19 | 19 | } | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -10,12 +10,12 @@ public static partial class tf | |||
| 10 | 10 | ||
| 11 | 11 | public static Tensor sub(Tensor a, Tensor b) => gen_math_ops.sub(a, b); | |
| 12 | 12 | ||
| 13 | - public static Tensor subtract<T>(Tensor x, T[] y, string name = "") where T : struct | ||
| 13 | + public static Tensor subtract<T>(Tensor x, T[] y, string name = null) where T : struct | ||
| 14 | 14 | => gen_math_ops.sub(x, ops.convert_to_tensor(y, dtype: x.dtype.as_base_dtype(), name: "y"), name); | |
| 15 | 15 | ||
| 16 | 16 | public static Tensor multiply(Tensor x, Tensor y) => gen_math_ops.mul(x, y); | |
| 17 | 17 | ||
| 18 | - public static Tensor divide<T>(Tensor x, T[] y, string name = "") where T : struct | ||
| 18 | + public static Tensor divide<T>(Tensor x, T[] y, string name = null) where T : struct | ||
| 19 | 19 | => x / ops.convert_to_tensor(y, dtype: x.dtype.as_base_dtype(), name: "y"); | |
| 20 | 20 | ||
| 21 | 21 | public static Tensor pow<T1, T2>(T1 x, T2 y) => gen_math_ops.pow(x, y); | |
@@ -28,7 +28,7 @@ public static Tensor divide<T>(Tensor x, T[] y, string name = "") where T : stru | |||
| 28 | 28 | /// <returns></returns> | |
| 29 | 29 | public static Tensor reduce_sum(Tensor input, int[] axis = null) => math_ops.reduce_sum(input); | |
| 30 | 30 | ||
| 31 | - public static Tensor cast(Tensor x, TF_DataType dtype = TF_DataType.DtInvalid, string name = "") | ||
| 31 | + public static Tensor cast(Tensor x, TF_DataType dtype = TF_DataType.DtInvalid, string name = null) | ||
| 32 | 32 | => math_ops.cast(x, dtype, name); | |
| 33 | 33 | } | |
| 34 | 34 | } | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -21,6 +21,6 @@ public static Tensor random_normal(int[] shape, | |||
| 21 | 21 | float stddev = 1.0f, | |
| 22 | 22 | TF_DataType dtype = TF_DataType.TF_FLOAT, | |
| 23 | 23 | int? seed = null, | |
| 24 | - string name = "") => random_ops.random_normal(shape, mean, stddev, dtype, seed, name); | ||
| 24 | + string name = null) => random_ops.random_normal(shape, mean, stddev, dtype, seed, name); | ||
| 25 | 25 | } | |
| 26 | 26 | } | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -6,7 +6,7 @@ namespace Tensorflow.Eager | |||
| 6 | 6 | { | |
| 7 | 7 | public class Execute | |
| 8 | 8 | { | |
| 9 | - public void record_gradient(string op_name, InputList inputs, Dictionary<string, object> attrs, Tensor[] results, string name = "") | ||
| 9 | + public void record_gradient(string op_name, InputList inputs, Dictionary<string, object> attrs, Tensor[] results, string name = null) | ||
| 10 | 10 | { | |
| 11 | 11 | pywrap_tfe_src.RecordGradient(op_name, inputs._inputs, attrs, results, name); | |
| 12 | 12 | } | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -10,7 +10,7 @@ namespace Tensorflow.Eager | |||
| 10 | 10 | /// </summary> | |
| 11 | 11 | public class pywrap_tfe_src | |
| 12 | 12 | { | |
| 13 | - public static void RecordGradient(string op_name, Tensor[] inputs, Dictionary<string, object> attrs, Tensor[] results, string name = "") | ||
| 13 | + public static void RecordGradient(string op_name, Tensor[] inputs, Dictionary<string, object> attrs, Tensor[] results, string name = null) | ||
| 14 | 14 | { | |
| 15 | 15 | var input_ids = inputs.Select(x => x.Id).ToArray(); | |
| 16 | 16 | var input_dtypes = inputs.Select(x => x.dtype).ToArray(); | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -12,7 +12,7 @@ public class importer | |||
| 12 | 12 | public static ITensorOrOperation[] import_graph_def(GraphDef graph_def, | |
| 13 | 13 | Dictionary<string, Tensor> input_map = null, | |
| 14 | 14 | string[] return_elements = null, | |
| 15 | - string name = "", | ||
| 15 | + string name = null, | ||
| 16 | 16 | OpList producer_op_list = null) | |
| 17 | 17 | { | |
| 18 | 18 | var op_dict = op_def_registry.get_registered_ops(); | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -123,7 +123,7 @@ public static (Dictionary<string, RefVariable>, ITensorOrOperation[]) import_sco | |||
| 123 | 123 | /// <param name="strip_default_attrs"></param> | |
| 124 | 124 | /// <param name="meta_info_def"></param> | |
| 125 | 125 | /// <returns></returns> | |
| 126 | - public static MetaGraphDef export_scoped_meta_graph(string filename = "", | ||
| 126 | + public static (MetaGraphDef, Dictionary<string, RefVariable>) export_scoped_meta_graph(string filename = "", | ||
| 127 | 127 | GraphDef graph_def = null, | |
| 128 | 128 | bool as_text = false, | |
| 129 | 129 | string unbound_inputs_col_name = "unbound_inputs", | |
@@ -138,7 +138,7 @@ public static MetaGraphDef export_scoped_meta_graph(string filename = "", | |||
| 138 | 138 | var var_list = new Dictionary<string, RefVariable>(); | |
| 139 | 139 | var variables = graph.get_collection(ops.GraphKeys.GLOBAL_VARIABLES); | |
| 140 | 140 | ||
| 141 | - foreach(var v in variables as RefVariable[]) | ||
| 141 | + foreach(var v in variables as List<RefVariable>) | ||
| 142 | 142 | { | |
| 143 | 143 | var_list[v.name] = v; | |
| 144 | 144 | } | |
@@ -151,15 +151,18 @@ public static MetaGraphDef export_scoped_meta_graph(string filename = "", | |||
| 151 | 151 | saver_def: saver_def, | |
| 152 | 152 | strip_default_attrs: strip_default_attrs); | |
| 153 | 153 | ||
| 154 | - throw new NotImplementedException("meta_graph.export_scoped_meta_graph"); | ||
| 154 | + if (!string.IsNullOrEmpty(filename)) | ||
| 155 | + graph_io.write_graph(scoped_meta_graph_def, "", filename, as_text: as_text); | ||
| 156 | + | ||
| 157 | + return (scoped_meta_graph_def, var_list); | ||
| 155 | 158 | } | |
| 156 | 159 | ||
| 157 | 160 | private static bool _should_include_node() | |
| 158 | 161 | { | |
| 159 | 162 | return true; | |
| 160 | 163 | } | |
| 161 | 164 | ||
| 162 | - private static byte[] create_meta_graph_def(MetaInfoDef meta_info_def = null, | ||
| 165 | + private static MetaGraphDef create_meta_graph_def(MetaInfoDef meta_info_def = null, | ||
| 163 | 166 | GraphDef graph_def = null, | |
| 164 | 167 | string export_scope = "", | |
| 165 | 168 | string exclude_nodes = "", | |
@@ -168,7 +171,7 @@ private static byte[] create_meta_graph_def(MetaInfoDef meta_info_def = null, | |||
| 168 | 171 | bool strip_default_attrs = false) | |
| 169 | 172 | { | |
| 170 | 173 | // Sets graph to default graph if it's not passed in. | |
| 171 | - var graph = ops.get_default_graph(); | ||
| 174 | + var graph = ops.get_default_graph().as_default(); | ||
| 172 | 175 | // Creates a MetaGraphDef proto. | |
| 173 | 176 | var meta_graph_def = new MetaGraphDef(); | |
| 174 | 177 | if (meta_info_def == null) | |
@@ -186,10 +189,55 @@ private static byte[] create_meta_graph_def(MetaInfoDef meta_info_def = null, | |||
| 186 | 189 | meta_graph_def.GraphDef = graph_def; | |
| 187 | 190 | ||
| 188 | 191 | // Fills in meta_info_def.stripped_op_list using the ops from graph_def. | |
| 189 | - if (meta_graph_def.MetaInfoDef.StrippedOpList.Op.Count == 0) | ||
| 192 | + if (meta_graph_def.MetaInfoDef.StrippedOpList == null || | ||
| 193 | + meta_graph_def.MetaInfoDef.StrippedOpList.Op.Count == 0) | ||
| 190 | 194 | meta_graph_def.MetaInfoDef.StrippedOpList = stripped_op_list_for_graph(meta_graph_def.GraphDef); | |
| 191 | 195 | ||
| 192 | - throw new NotImplementedException("create_meta_graph_def"); | ||
| 196 | + var clist = graph.get_all_collection_keys(); | ||
| 197 | + foreach(var ctype in clist) | ||
| 198 | + { | ||
| 199 | + if (clear_extraneous_savers) | ||
| 200 | + { | ||
| 201 | + throw new NotImplementedException("create_meta_graph_def clear_extraneous_savers"); | ||
| 202 | + } | ||
| 203 | + else | ||
| 204 | + { | ||
| 205 | + add_collection_def(meta_graph_def, ctype, graph); | ||
| 206 | + } | ||
| 207 | + } | ||
| 208 | + | ||
| 209 | + return meta_graph_def; | ||
| 210 | + } | ||
| 211 | + | ||
| 212 | + private static void add_collection_def(MetaGraphDef meta_graph_def, | ||
| 213 | + string key, | ||
| 214 | + Graph graph = null, | ||
| 215 | + string export_scope = "") | ||
| 216 | + { | ||
| 217 | + if (!meta_graph_def.CollectionDef.ContainsKey(key)) | ||
| 218 | + meta_graph_def.CollectionDef[key] = new CollectionDef(); | ||
| 219 | + var col_def = meta_graph_def.CollectionDef[key]; | ||
| 220 | + | ||
| 221 | + switch (graph.get_collection(key)) | ||
| 222 | + { | ||
| 223 | + case List<RefVariable> collection_list: | ||
| 224 | + col_def.BytesList = new Types.BytesList(); | ||
| 225 | + foreach (var x in collection_list) | ||
| 226 | + { | ||
| 227 | + var proto = x.to_proto(export_scope); | ||
| 228 | + col_def.BytesList.Value.Add(proto.ToByteString()); | ||
| 229 | + } | ||
| 230 | + | ||
| 231 | + break; | ||
| 232 | + case List<object> collection_list: | ||
| 233 | + col_def.NodeList = new Types.NodeList(); | ||
| 234 | + foreach (var x in collection_list) | ||
| 235 | + if (x is ITensorOrOperation x2) | ||
| 236 | + col_def.NodeList.Value.Add(ops.strip_name_scope(x2.name, export_scope)); | ||
| 237 | + break; | ||
| 238 | + case List<Operation> collection_list: | ||
| 239 | + break; | ||
| 240 | + } | ||
| 193 | 241 | } | |
| 194 | 242 | ||
| 195 | 243 | private static OpList stripped_op_list_for_graph(GraphDef graph_def) | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -118,7 +118,7 @@ private ITensorOrOperation _as_graph_element_locked(object obj, bool allow_tenso | |||
| 118 | 118 | ||
| 119 | 119 | if (obj is Tensor tensor && allow_tensor) | |
| 120 | 120 | { | |
| 121 | - if (tensor.Graph.Equals(this)) | ||
| 121 | + if (tensor.graph.Equals(this)) | ||
| 122 | 122 | { | |
| 123 | 123 | return tensor; | |
| 124 | 124 | } | |
@@ -164,7 +164,7 @@ private void _check_not_finalized() | |||
| 164 | 164 | } | |
| 165 | 165 | ||
| 166 | 166 | public unsafe Operation create_op(string op_type, Tensor[] inputs, TF_DataType[] dtypes, | |
| 167 | - TF_DataType[] input_types = null, string name = "", | ||
| 167 | + TF_DataType[] input_types = null, string name = null, | ||
| 168 | 168 | Dictionary<string, AttrValue> attrs = null, OpDef op_def = null) | |
| 169 | 169 | { | |
| 170 | 170 | if (inputs == null) | |
| Back | FazBrowse Home | New Git URL |
0 commit comments