| FazBrowse GitHub Viewer | Trending | | Home |
| Tools: [Download Repo ZIP] [Original HTTPS Page] |
1 parent db8c088 commit 879067d
12 files changed
| 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.Diagnostics; | ||
| 3 | 4 | using System.Text; | |
| 4 | 5 | using NumSharp; | |
| 5 | 6 | using Tensorflow; | |
@@ -21,7 +22,12 @@ public MnistDataSet(NDArray images, NDArray labels, Type dataType, bool reshape) | |||
| 21 | 22 | ||
| 22 | 23 | images = images.reshape(images.shape[0], images.shape[1] * images.shape[2]); | |
| 23 | 24 | images.astype(dataType); | |
| 25 | + // for debug np.multiply performance | ||
| 26 | + var sw = new Stopwatch(); | ||
| 27 | + sw.Start(); | ||
| 24 | 28 | images = np.multiply(images, 1.0f / 255.0f); | |
| 29 | + sw.Stop(); | ||
| 30 | + Console.WriteLine($"{sw.ElapsedMilliseconds}ms"); | ||
| 25 | 31 | Data = images; | |
| 26 | 32 | ||
| 27 | 33 | labels.astype(dataType); | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -14,10 +14,29 @@ You may obtain a copy of the License at | |||
| 14 | 14 | limitations under the License. | |
| 15 | 15 | ******************************************************************************/ | |
| 16 | 16 | ||
| 17 | + using System; | ||
| 18 | + | ||
| 17 | 19 | namespace Tensorflow | |
| 18 | 20 | { | |
| 19 | 21 | public static partial class tf | |
| 20 | 22 | { | |
| 23 | + public static Tensor while_loop(Func<Tensor, Tensor> cond, Func<Tensor, Tensor> body, Tensor[] loop_vars, | ||
| 24 | + TensorShape shape_invariants = null, | ||
| 25 | + int parallel_iterations = 10, | ||
| 26 | + bool back_prop = true, | ||
| 27 | + bool swap_memory = false, | ||
| 28 | + string name = null, | ||
| 29 | + int? maximum_iterations = null, | ||
| 30 | + bool return_same_structure = false) | ||
| 31 | + => control_flow_ops.while_loop(cond, body, loop_vars, | ||
| 32 | + shape_invariants: shape_invariants, | ||
| 33 | + parallel_iterations: parallel_iterations, | ||
| 34 | + back_prop: back_prop, | ||
| 35 | + swap_memory: swap_memory, | ||
| 36 | + name: name, | ||
| 37 | + maximum_iterations: maximum_iterations, | ||
| 38 | + return_same_structure: return_same_structure); | ||
| 39 | + | ||
| 21 | 40 | public static _ControlDependenciesController control_dependencies(Operation[] control_inputs) | |
| 22 | 41 | => ops.control_dependencies(control_inputs); | |
| 23 | 42 | } | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -39,8 +39,8 @@ public static Tensor acos(Tensor x, string name = null) | |||
| 39 | 39 | public static Tensor asin(Tensor x, string name = null) | |
| 40 | 40 | => gen_math_ops.asin(x, name); | |
| 41 | 41 | ||
| 42 | - public static Tensor add<Tx, Ty>(Tx a, Ty b) | ||
| 43 | - => gen_math_ops.add(a, b); | ||
| 42 | + public static Tensor add<Tx, Ty>(Tx a, Ty b, string name = null) | ||
| 43 | + => gen_math_ops.add(a, b, name: name); | ||
| 44 | 44 | ||
| 45 | 45 | /// <summary> | |
| 46 | 46 | /// Computes atan of x element-wise. | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -33,7 +33,7 @@ public partial class Graph | |||
| 33 | 33 | /// </summary> | |
| 34 | 34 | /// <param name="input_ops">The data input ops for an op to be created.</param> | |
| 35 | 35 | /// <returns>A list of control inputs for the op to be created.</returns> | |
| 36 | - private ITensorOrOperation[] _control_dependencies_for_inputs(ITensorOrOperation[] input_ops) | ||
| 36 | + public ITensorOrOperation[] _control_dependencies_for_inputs(ITensorOrOperation[] input_ops) | ||
| 37 | 37 | { | |
| 38 | 38 | var ret = new List<ITensorOrOperation>(); | |
| 39 | 39 | ||
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -53,6 +53,11 @@ public Tensor pivot | |||
| 53 | 53 | protected Stack<ControlFlowContext> _context_stack; | |
| 54 | 54 | protected ControlFlowContext _outer_context; | |
| 55 | 55 | ||
| 56 | + /// <summary> | ||
| 57 | + /// The keys are the names of tensors referenced by but external to this | ||
| 58 | + /// context. Each value is the Tensor that should be used by this context to | ||
| 59 | + /// access the key value (e.g. a switch output guarding a cond input value). | ||
| 60 | + /// </summary> | ||
| 56 | 61 | protected Dictionary<string, ITensorOrOperation> _external_values; | |
| 57 | 62 | ||
| 58 | 63 | public ControlFlowContext() | |
@@ -68,6 +73,12 @@ public void __init__(ValuesDef values_def = null, string import_scope = null) | |||
| 68 | 73 | _outer_context = ops.get_default_graph()._get_control_flow_context(); | |
| 69 | 74 | if (values_def != null) | |
| 70 | 75 | _init_values_from_proto(values_def, import_scope: import_scope); | |
| 76 | + else | ||
| 77 | + { | ||
| 78 | + _values = new HashSet<string>(); | ||
| 79 | + _external_values = new Dictionary<string, ITensorOrOperation>(); | ||
| 80 | + } | ||
| 81 | + | ||
| 71 | 82 | } | |
| 72 | 83 | ||
| 73 | 84 | public void __enter__() | |
@@ -114,6 +125,27 @@ public virtual void Enter() | |||
| 114 | 125 | graph._set_control_flow_context(this); | |
| 115 | 126 | } | |
| 116 | 127 | ||
| 128 | + protected virtual Tensor _Enter(Tensor data, string frame_name, | ||
| 129 | + bool is_constant = false, | ||
| 130 | + int parallel_iterations = 10, | ||
| 131 | + bool use_ref = true, | ||
| 132 | + bool use_input_shape = true, | ||
| 133 | + string name = null) | ||
| 134 | + { | ||
| 135 | + Tensor result; | ||
| 136 | + data = ops.internal_convert_to_tensor_or_indexed_slices(data, as_ref: true); | ||
| 137 | + if (data.dtype.is_ref_dtype() && use_ref) | ||
| 138 | + throw new NotImplementedException("_Enter"); | ||
| 139 | + else | ||
| 140 | + result = gen_control_flow_ops.enter( | ||
| 141 | + data, frame_name, is_constant, parallel_iterations, name: name); | ||
| 142 | + | ||
| 143 | + if (use_input_shape) | ||
| 144 | + result.SetShape(data.TensorShape); | ||
| 145 | + | ||
| 146 | + return result; | ||
| 147 | + } | ||
| 148 | + | ||
| 117 | 149 | /// <summary> | |
| 118 | 150 | /// Exit this control flow context. | |
| 119 | 151 | /// </summary> | |
@@ -184,6 +216,10 @@ public static bool IsContainingContext(ControlFlowContext ctxt, ControlFlowConte | |||
| 184 | 216 | return true; | |
| 185 | 217 | } | |
| 186 | 218 | ||
| 219 | + protected virtual bool _IsInOuterContext(Operation op) | ||
| 220 | + { | ||
| 221 | + throw new NotImplementedException("_IsInOuterContext"); | ||
| 222 | + } | ||
| 187 | 223 | ||
| 188 | 224 | protected virtual void _RemoveExternalControlEdges(Operation op) | |
| 189 | 225 | { | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -15,8 +15,12 @@ limitations under the License. | |||
| 15 | 15 | ******************************************************************************/ | |
| 16 | 16 | ||
| 17 | 17 | using System; | |
| 18 | + using System.Collections.Generic; | ||
| 19 | + using System.Linq; | ||
| 18 | 20 | using Tensorflow.Operations.ControlFlows; | |
| 21 | + using Tensorflow.Util; | ||
| 19 | 22 | using static Tensorflow.Python; | |
| 23 | + using static Tensorflow.control_flow_ops; | ||
| 20 | 24 | ||
| 21 | 25 | namespace Tensorflow.Operations | |
| 22 | 26 | { | |
@@ -32,10 +36,14 @@ public class WhileContext : ControlFlowContext | |||
| 32 | 36 | bool _swap_memory; | |
| 33 | 37 | Tensor _pivot_for_pred; | |
| 34 | 38 | Tensor _pivot_for_body; | |
| 35 | - Tensor[] _loop_exits; | ||
| 36 | - Tensor[] _loop_enters; | ||
| 39 | + List<Tensor> _loop_exits; | ||
| 40 | + List<Tensor> _loop_enters; | ||
| 41 | + Graph _graph; | ||
| 42 | + public override GradLoopState grad_state => _grad_state; | ||
| 43 | + public override bool back_prop => _back_prop; | ||
| 37 | 44 | ||
| 38 | - public WhileContext(int parallel_iterations = 10, | ||
| 45 | + public WhileContext(int? maximum_iterations = null, | ||
| 46 | + int parallel_iterations = 10, | ||
| 39 | 47 | bool back_prop = true, | |
| 40 | 48 | bool swap_memory = false, | |
| 41 | 49 | string name = "while_context", | |
@@ -49,12 +57,27 @@ public WhileContext(int parallel_iterations = 10, | |||
| 49 | 57 | } | |
| 50 | 58 | else | |
| 51 | 59 | { | |
| 52 | - | ||
| 60 | + __init__(); | ||
| 61 | + _init_from_args(maximum_iterations, parallel_iterations, back_prop, swap_memory, name); | ||
| 53 | 62 | } | |
| 54 | 63 | ||
| 55 | 64 | _grad_state = grad_state; | |
| 56 | 65 | } | |
| 57 | 66 | ||
| 67 | + private void _init_from_args(int? maximum_iterations, | ||
| 68 | + int parallel_iterations, | ||
| 69 | + bool back_prop, | ||
| 70 | + bool swap_memory, | ||
| 71 | + string name) | ||
| 72 | + { | ||
| 73 | + _name = ops.get_default_graph().unique_name(name); | ||
| 74 | + _back_prop = back_prop; | ||
| 75 | + _swap_memory = swap_memory; | ||
| 76 | + _loop_exits = new List<Tensor>(); | ||
| 77 | + _loop_enters = new List<Tensor>(); | ||
| 78 | + _graph = ops.get_default_graph(); | ||
| 79 | + } | ||
| 80 | + | ||
| 58 | 81 | private void _init_from_proto(WhileContextDef context_def, string import_scope = null) | |
| 59 | 82 | { | |
| 60 | 83 | var g = ops.get_default_graph(); | |
@@ -70,26 +93,156 @@ private void _init_from_proto(WhileContextDef context_def, string import_scope = | |||
| 70 | 93 | // The boolean tensor for loop termination condition. | |
| 71 | 94 | _pivot = g.as_graph_element(ops.prepend_name_scope(context_def.PivotName, import_scope)) as Tensor; | |
| 72 | 95 | // The list of exit tensors for loop variables. | |
| 73 | - _loop_exits = new Tensor[context_def.LoopExitNames.Count]; | ||
| 96 | + _loop_exits = new List<Tensor>(); | ||
| 74 | 97 | foreach (var (i, exit_name) in enumerate(context_def.LoopExitNames)) | |
| 75 | - _loop_exits[i] = g.as_graph_element(ops.prepend_name_scope(exit_name, import_scope)) as Tensor; | ||
| 98 | + _loop_exits.Add(g.as_graph_element(ops.prepend_name_scope(exit_name, import_scope)) as Tensor); | ||
| 76 | 99 | // The list of enter tensors for loop variables. | |
| 77 | - _loop_enters = new Tensor[context_def.LoopEnterNames.Count]; | ||
| 100 | + _loop_enters = new List<Tensor>(); | ||
| 78 | 101 | foreach (var (i, enter_name) in enumerate(context_def.LoopEnterNames)) | |
| 79 | - _loop_enters[i] = g.as_graph_element(ops.prepend_name_scope(enter_name, import_scope)) as Tensor; | ||
| 102 | + _loop_enters.Add(g.as_graph_element(ops.prepend_name_scope(enter_name, import_scope)) as Tensor); | ||
| 80 | 103 | ||
| 81 | 104 | __init__(values_def: context_def.ValuesDef, import_scope: import_scope); | |
| 82 | 105 | } | |
| 83 | 106 | ||
| 84 | - public override WhileContext GetWhileContext() | ||
| 107 | + /// <summary> | ||
| 108 | + /// Add the loop termination condition and body to the graph. | ||
| 109 | + /// </summary> | ||
| 110 | + public Tensor[] BuildLoop(Func<Tensor, Tensor> pred, | ||
| 111 | + Func<Tensor, Tensor> body, | ||
| 112 | + Tensor[] loop_vars, | ||
| 113 | + TensorShape shape_invariants, | ||
| 114 | + bool return_same_structure) | ||
| 85 | 115 | { | |
| 86 | - return this; | ||
| 116 | + // Keep original_loop_vars to identify which are TensorArrays | ||
| 117 | + var original_loop_vars = loop_vars; | ||
| 118 | + // Convert TensorArrays to their flow variables | ||
| 119 | + Enter(); | ||
| 120 | + var(original_body_result, exit_vars) = _BuildLoop( | ||
| 121 | + pred, body, original_loop_vars, loop_vars, shape_invariants); | ||
| 122 | + Exit(); | ||
| 123 | + | ||
| 124 | + var flat_result = original_body_result; | ||
| 125 | + | ||
| 126 | + var exit_vars_with_tensor_arrays = _convert_flows_to_tensorarrays(flat_result, exit_vars); | ||
| 127 | + var packed_exit_vars = nest.pack_sequence_as( | ||
| 128 | + structure: original_body_result, | ||
| 129 | + flat_sequence: exit_vars_with_tensor_arrays); | ||
| 130 | + | ||
| 131 | + return packed_exit_vars as Tensor[]; | ||
| 87 | 132 | } | |
| 88 | 133 | ||
| 134 | + private (Tensor[], Tensor[]) _BuildLoop(Func<Tensor, Tensor> pred, | ||
| 135 | + Func<Tensor, Tensor> body, | ||
| 136 | + Tensor[] original_loop_vars, | ||
| 137 | + Tensor[] loop_vars, | ||
| 138 | + TensorShape shape_invariants) | ||
| 139 | + { | ||
| 140 | + var flat_loop_vars = original_loop_vars; | ||
| 89 | 141 | ||
| 90 | - public override GradLoopState grad_state => _grad_state; | ||
| 142 | + // Let the context know the loop variables so the loop variables | ||
| 143 | + // would be added in the outer contexts properly. | ||
| 144 | + _InitializeValues(loop_vars); | ||
| 145 | + var real_vars = loop_vars; | ||
| 146 | + Tensor[] enter_vars = null; | ||
| 147 | + tf_with(ops.control_dependencies(null), delegate | ||
| 148 | + { | ||
| 149 | + enter_vars = real_vars.Select(x => _Enter(x, | ||
| 150 | + _name, | ||
| 151 | + is_constant: false, | ||
| 152 | + parallel_iterations: _parallel_iterations, | ||
| 153 | + use_input_shape: shape_invariants == null)) | ||
| 154 | + .ToArray(); | ||
| 91 | 155 | ||
| 92 | - public override bool back_prop => _back_prop; | ||
| 156 | + foreach(var x in enter_vars) | ||
| 157 | + { | ||
| 158 | + x.graph.prevent_feeding(x); | ||
| 159 | + if (_outer_context != null) | ||
| 160 | + _outer_context.AddInnerOp(x.op); | ||
| 161 | + } | ||
| 162 | + }); | ||
| 163 | + | ||
| 164 | + // Finds the closest enclosing non-None control pivot. | ||
| 165 | + var outer_context = _outer_context; | ||
| 166 | + while (outer_context != null) | ||
| 167 | + { | ||
| 168 | + | ||
| 169 | + } | ||
| 170 | + | ||
| 171 | + _SetShapeInvariants(real_vars, enter_vars, shape_invariants); | ||
| 172 | + | ||
| 173 | + // Fix the control inputs and control flow context of these enter ops. | ||
| 174 | + _FixControlInputsAndContext(enter_vars); | ||
| 175 | + _InitializeValues(enter_vars); | ||
| 176 | + _loop_enters = enter_vars.ToList(); | ||
| 177 | + | ||
| 178 | + var merge_vars = enter_vars | ||
| 179 | + .Select(x => merge(new[] { x, x })) | ||
| 180 | + .ToArray(); | ||
| 181 | + | ||
| 182 | + _pivot_for_pred = merge_vars[0]; | ||
| 183 | + | ||
| 184 | + // Build the graph for pred. | ||
| 185 | + var merge_vars_with_tensor_arrays = _convert_flows_to_tensorarrays(flat_loop_vars, merge_vars); | ||
| 186 | + // var packed_vars = nest.pack_sequence_as(original_loop_vars, merge_vars_with_tensor_arrays); | ||
| 187 | + var c = ops.convert_to_tensor(pred(merge_vars_with_tensor_arrays[0])); | ||
| 188 | + _pivot = gen_control_flow_ops.loop_cond(c, name: "LoopCond"); | ||
| 189 | + var switch_vars = merge_vars.Select(x => _SwitchRefOrTensor(x, _pivot)) | ||
| 190 | + .ToArray(); | ||
| 191 | + | ||
| 192 | + // Build the graph for body. | ||
| 193 | + var vars_for_body = switch_vars.Select(x => _Identity(x[1])).ToArray(); | ||
| 194 | + // Convert TensorArray flow variables inside the context back into | ||
| 195 | + // their associated TensorArrays for calling the body. | ||
| 196 | + var packed_vars_for_body = _convert_flows_to_tensorarrays(flat_loop_vars, vars_for_body); | ||
| 197 | + var body_result = body(packed_vars_for_body[0]); | ||
| 198 | + var post_summaries = ops.get_collection(ops.GraphKeys._SUMMARY_COLLECTION); | ||
| 199 | + | ||
| 200 | + // Store body_result to keep track of TensorArrays returned by body | ||
| 201 | + var original_body_result = new[] { body_result }; | ||
| 202 | + // Convert TensorArrays returned by body into their flow variables | ||
| 203 | + var result = new[] { body_result }; | ||
| 204 | + | ||
| 205 | + var next_vars = new List<Tensor>(); | ||
| 206 | + foreach (var (m, v) in zip(merge_vars, result)) | ||
| 207 | + next_vars.Add(_AddNextAndBackEdge(m, v)); | ||
| 208 | + | ||
| 209 | + // Add the exit ops. | ||
| 210 | + var exit_vars = switch_vars.Select(x => exit(x[0])).ToList(); | ||
| 211 | + _loop_exits = exit_vars; | ||
| 212 | + | ||
| 213 | + // Exit the loop. | ||
| 214 | + // ExitResult(exit_vars); | ||
| 215 | + return (original_body_result, exit_vars.ToArray()); | ||
| 216 | + } | ||
| 217 | + | ||
| 218 | + private void _FixControlInputsAndContext(Tensor[] enters) | ||
| 219 | + { | ||
| 220 | + var graph = ops.get_default_graph(); | ||
| 221 | + foreach(var e in enters) | ||
| 222 | + { | ||
| 223 | + var inp_op = e.op.inputs[0].op; | ||
| 224 | + var control_inputs = graph._control_dependencies_for_inputs(new[] { inp_op }); | ||
| 225 | + // op for op in control_inputs if self._IsInOuterContext(op) | ||
| 226 | + var outer_control_inputs = control_inputs.Where(x => _IsInOuterContext(x.op)) | ||
| 227 | + .Select(x => x.op) | ||
| 228 | + .ToArray(); | ||
| 229 | + e.op._set_control_flow_context(this); | ||
| 230 | + e.op._add_control_inputs(outer_control_inputs); | ||
| 231 | + graph._record_op_seen_by_control_dependencies(e.op); | ||
| 232 | + } | ||
| 233 | + } | ||
| 234 | + | ||
| 235 | + private void _InitializeValues(Tensor[] values) | ||
| 236 | + { | ||
| 237 | + _values = new HashSet<string>(); | ||
| 238 | + foreach(var x in values) | ||
| 239 | + _values.Add(x.name); | ||
| 240 | + } | ||
| 241 | + | ||
| 242 | + public override WhileContext GetWhileContext() | ||
| 243 | + { | ||
| 244 | + return this; | ||
| 245 | + } | ||
| 93 | 246 | ||
| 94 | 247 | public WhileContext from_proto(WhileContextDef proto, string import_scope) | |
| 95 | 248 | { | |
| Back | FazBrowse Home | New Git URL |
0 commit comments