| FazBrowse GitHub Viewer | Trending | | Home |
| Tools: [Download Repo ZIP] [Original HTTPS Page] |
17 files changed
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -27,7 +27,7 @@ public static Tensor acos(Tensor x, string name = null) | |||
| 27 | 27 | public static Tensor asin(Tensor x, string name = null) | |
| 28 | 28 | => gen_math_ops.asin(x, name); | |
| 29 | 29 | ||
| 30 | - public static Tensor add(Tensor a, Tensor b) | ||
| 30 | + public static Tensor add<Tx, Ty>(Tx a, Ty b) | ||
| 31 | 31 | => gen_math_ops.add(a, b); | |
| 32 | 32 | ||
| 33 | 33 | /// <summary> | |
@@ -251,7 +251,7 @@ public static Tensor maximum<T1, T2>(T1 x, T2 y, string name = null) | |||
| 251 | 251 | public static Tensor minimum<T1, T2>(T1 x, T2 y, string name = null) | |
| 252 | 252 | => gen_math_ops.minimum(x, y, name: name); | |
| 253 | 253 | ||
| 254 | - public static Tensor multiply(Tensor x, Tensor y) | ||
| 254 | + public static Tensor multiply<Tx, Ty>(Tx x, Ty y) | ||
| 255 | 255 | => gen_math_ops.mul(x, y); | |
| 256 | 256 | ||
| 257 | 257 | public static Tensor negative(Tensor x, string name = null) | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -4,6 +4,7 @@ | |||
| 4 | 4 | using System.IO; | |
| 5 | 5 | using System.Linq; | |
| 6 | 6 | using System.Text; | |
| 7 | + using Tensorflow.Operations; | ||
| 7 | 8 | using static Tensorflow.CollectionDef; | |
| 8 | 9 | using static Tensorflow.MetaGraphDef.Types; | |
| 9 | 10 | ||
@@ -95,15 +96,29 @@ public static (Dictionary<string, RefVariable>, ITensorOrOperation[]) import_sco | |||
| 95 | 96 | } | |
| 96 | 97 | else | |
| 97 | 98 | { | |
| 98 | - throw new NotImplementedException("import_scoped_meta_graph_with_return_elements"); | ||
| 99 | + foreach(var value in col.Value.BytesList.Value) | ||
| 100 | + { | ||
| 101 | + switch (col.Key) | ||
| 102 | + { | ||
| 103 | + case "cond_context": | ||
| 104 | + var proto = CondContextDef.Parser.ParseFrom(value); | ||
| 105 | + var condContext = new CondContext().from_proto(proto, import_scope); | ||
| 106 | + graph.add_to_collection(col.Key, condContext); | ||
| 107 | + break; | ||
| 108 | + default: | ||
| 109 | + throw new NotImplementedException("import_scoped_meta_graph_with_return_elements"); | ||
| 110 | + } | ||
| 111 | + } | ||
| 99 | 112 | } | |
| 100 | 113 | ||
| 101 | 114 | break; | |
| 115 | + default: | ||
| 116 | + throw new NotImplementedException("import_scoped_meta_graph_with_return_elements"); | ||
| 102 | 117 | } | |
| 103 | 118 | } | |
| 104 | 119 | ||
| 105 | - var variables = graph.get_collection(ops.GraphKeys.GLOBAL_VARIABLES, | ||
| 106 | - scope: scope_to_prepend_to_names) as List<RefVariable>; | ||
| 120 | + var variables = graph.get_collection<RefVariable>(ops.GraphKeys.GLOBAL_VARIABLES, | ||
| 121 | + scope: scope_to_prepend_to_names); | ||
| 107 | 122 | var var_list = new Dictionary<string, RefVariable>(); | |
| 108 | 123 | variables.ForEach(v => var_list[ops.strip_name_scope(v.name, scope_to_prepend_to_names)] = v); | |
| 109 | 124 | ||
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -412,6 +412,11 @@ public object get_collection(string name, string scope = null) | |||
| 412 | 412 | return _collections.ContainsKey(name) ? _collections[name] : null; | |
| 413 | 413 | } | |
| 414 | 414 | ||
| 415 | + public List<T> get_collection<T>(string name, string scope = null) | ||
| 416 | + { | ||
| 417 | + return _collections.ContainsKey(name) ? _collections[name] as List<T> : new List<T>(); | ||
| 418 | + } | ||
| 419 | + | ||
| 415 | 420 | public object get_collection_ref(string name) | |
| 416 | 421 | { | |
| 417 | 422 | if (!_collections.ContainsKey(name)) | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -8,7 +8,7 @@ namespace Tensorflow.Operations | |||
| 8 | 8 | /// <summary> | |
| 9 | 9 | /// The context for the conditional construct. | |
| 10 | 10 | /// </summary> | |
| 11 | - public class CondContext : ControlFlowContext | ||
| 11 | + public class CondContext : ControlFlowContext, IProtoBuf<CondContextDef, CondContext> | ||
| 12 | 12 | { | |
| 13 | 13 | ||
| 14 | 14 | ||
@@ -35,16 +35,20 @@ public class CondContext : ControlFlowContext | |||
| 35 | 35 | /// <param name="name">Name of the `CondContext` python object.</param> | |
| 36 | 36 | /// <param name="context_def"></param> | |
| 37 | 37 | /// <param name="import_scope"></param> | |
| 38 | - public CondContext(Tensor pred, | ||
| 39 | - Tensor pivot, | ||
| 40 | - int branch, | ||
| 38 | + public CondContext(Tensor pred = null, | ||
| 39 | + Tensor pivot = null, | ||
| 40 | + int? branch = null, | ||
| 41 | 41 | string name = "cond_text", | |
| 42 | - object context_def = null, | ||
| 42 | + CondContextDef context_def = null, | ||
| 43 | 43 | string import_scope = null) | |
| 44 | 44 | { | |
| 45 | + if (pred == null && context_def == null) return; | ||
| 46 | + | ||
| 45 | 47 | _name = ops.get_default_graph().unique_name(name); | |
| 46 | - if (context_def != null) | ||
| 47 | - throw new NotImplementedException("CondContext context_def is not null"); | ||
| 48 | + if (context_def != null) | ||
| 49 | + { | ||
| 50 | + _init_from_proto(context_def, import_scope: import_scope); | ||
| 51 | + } | ||
| 48 | 52 | else | |
| 49 | 53 | { | |
| 50 | 54 | // Initializes the default fields. | |
@@ -61,6 +65,18 @@ public CondContext(Tensor pred, | |||
| 61 | 65 | } | |
| 62 | 66 | } | |
| 63 | 67 | ||
| 68 | + private void _init_from_proto(CondContextDef context_def, string import_scope = null) | ||
| 69 | + { | ||
| 70 | + var g = ops.get_default_graph(); | ||
| 71 | + _name = ops.prepend_name_scope(context_def.ContextName, import_scope); | ||
| 72 | + var p1 = ops.prepend_name_scope(context_def.PredName, import_scope); | ||
| 73 | + _pred = g.as_graph_element(p1) as Tensor; | ||
| 74 | + var p2 = ops.prepend_name_scope(context_def.PivotName, import_scope); | ||
| 75 | + _pivot = g.as_graph_element(p2) as Tensor; | ||
| 76 | + _branch = context_def.Branch; | ||
| 77 | + __init__(values_def: context_def.ValuesDef, import_scope: import_scope); | ||
| 78 | + } | ||
| 79 | + | ||
| 64 | 80 | /// <summary> | |
| 65 | 81 | /// Add `val` to the current context and its outer context recursively. | |
| 66 | 82 | /// </summary> | |
@@ -230,6 +246,22 @@ private Tensor _ProcessOutputTensor(Tensor val) | |||
| 230 | 246 | public override void AddInnerOp(Operation resultOp) | |
| 231 | 247 | { | |
| 232 | 248 | throw new NotImplementedException(); | |
| 233 | - } | ||
| 249 | + } | ||
| 250 | + | ||
| 251 | + public CondContextDef to_proto(string export_scope) | ||
| 252 | + { | ||
| 253 | + throw new NotImplementedException(); | ||
| 254 | + } | ||
| 255 | + | ||
| 256 | + public CondContext from_proto(CondContextDef proto, string import_scope) | ||
| 257 | + { | ||
| 258 | + var ret = new CondContext(context_def: proto, import_scope: import_scope); | ||
| 259 | + | ||
| 260 | + ret.Enter(); | ||
| 261 | + foreach (var nested_def in proto.NestedContexts) | ||
| 262 | + throw new NotImplementedException(""); | ||
| 263 | + ret.Exit(); | ||
| 264 | + return ret; | ||
| 265 | + } | ||
| 234 | 266 | } | |
| 235 | 267 | } | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -32,6 +32,8 @@ public abstract class ControlFlowContext : Python, IPython, IControlFlowContext | |||
| 32 | 32 | protected Stack<IControlFlowContext> _context_stack; | |
| 33 | 33 | protected IControlFlowContext _outer_context; | |
| 34 | 34 | ||
| 35 | + protected Dictionary<string, ITensorOrOperation> _external_values; | ||
| 36 | + | ||
| 35 | 37 | public ControlFlowContext() | |
| 36 | 38 | { | |
| 37 | 39 | _context_stack = new Stack<IControlFlowContext>(); | |
@@ -40,15 +42,43 @@ public ControlFlowContext() | |||
| 40 | 42 | public string name { get => _name; } | |
| 41 | 43 | protected string _name; | |
| 42 | 44 | ||
| 43 | - public void __init__() | ||
| 45 | + public void __init__(ValuesDef values_def = null, string import_scope = null) | ||
| 44 | 46 | { | |
| 45 | - | ||
| 47 | + _outer_context = ops.get_default_graph()._get_control_flow_context(); | ||
| 48 | + if (values_def != null) | ||
| 49 | + _init_values_from_proto(values_def, import_scope: import_scope); | ||
| 46 | 50 | } | |
| 47 | 51 | ||
| 48 | 52 | public void __enter__() | |
| 49 | 53 | { | |
| 50 | 54 | } | |
| 51 | 55 | ||
| 56 | + /// <summary> | ||
| 57 | + /// Initializes values and external_values from `ValuesDef` protocol buffer. | ||
| 58 | + /// </summary> | ||
| 59 | + /// <param name="values_def"></param> | ||
| 60 | + /// <param name="import_scope"></param> | ||
| 61 | + protected void _init_values_from_proto(ValuesDef values_def, string import_scope = null) | ||
| 62 | + { | ||
| 63 | + _external_values = new Dictionary<string, ITensorOrOperation>(); | ||
| 64 | + foreach (var value in values_def.Values) | ||
| 65 | + _values.Add(value); | ||
| 66 | + var g = ops.get_default_graph(); | ||
| 67 | + foreach(var value in values_def.ExternalValues) | ||
| 68 | + { | ||
| 69 | + var k = ops.prepend_name_scope(value.Key, import_scope); | ||
| 70 | + var v = value.Value; | ||
| 71 | + _external_values[k] = g.as_graph_element(ops.prepend_name_scope(v, import_scope)); | ||
| 72 | + } | ||
| 73 | + | ||
| 74 | + var op_names = _values.Where(x => !_external_values.ContainsKey(x)) | ||
| 75 | + .Select(x => x.Split(':')[0]) | ||
| 76 | + .ToArray(); | ||
| 77 | + | ||
| 78 | + foreach (var op in op_names) | ||
| 79 | + (g.as_graph_element(op) as Operation)._set_control_flow_context(this); | ||
| 80 | + } | ||
| 81 | + | ||
| 52 | 82 | public void __exit__() | |
| 53 | 83 | { | |
| 54 | 84 | } | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -42,8 +42,8 @@ public unsafe Operation[] GetControlOutputs() | |||
| 42 | 42 | if (NumControlOutputs > 0) | |
| 43 | 43 | { | |
| 44 | 44 | IntPtr control_output_handle = Marshal.AllocHGlobal(Marshal.SizeOf<IntPtr>() * NumControlOutputs); | |
| 45 | - c_api.TF_OperationGetControlOutputs(_handle, control_output_handle, NumControlInputs); | ||
| 46 | - for (int i = 0; i < NumControlInputs; i++) | ||
| 45 | + c_api.TF_OperationGetControlOutputs(_handle, control_output_handle, NumControlOutputs); | ||
| 46 | + for (int i = 0; i < NumControlOutputs; i++) | ||
| 47 | 47 | { | |
| 48 | 48 | var handle = control_output_handle + Marshal.SizeOf<IntPtr>() * i; | |
| 49 | 49 | control_outputs[i] = new Operation(*(IntPtr*)handle); | |
| Back | FazBrowse Home | New Git URL |
0 commit comments