| FazBrowse GitHub Viewer | Trending | | Home |
| Tools: [Download Repo ZIP] [Original HTTPS Page] |
1 parent 43273b3 commit 444cc42
6 files changed
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -36,7 +36,7 @@ public unsafe TF_Input[] OutputConsumers(int index, int max_consumers) | |||
| 36 | 36 | var handle = Marshal.AllocHGlobal(size); | |
| 37 | 37 | int num = c_api.TF_OperationOutputConsumers(new TF_Output(_handle, index), handle, max_consumers); | |
| 38 | 38 | var consumers = new TF_Input[num]; | |
| 39 | - for(int i = 0; i < num; i++) | ||
| 39 | + for (int i = 0; i < num; i++) | ||
| 40 | 40 | { | |
| 41 | 41 | consumers[i] = Marshal.PtrToStructure<TF_Input>(handle + i * size); | |
| 42 | 42 | } | |
@@ -50,7 +50,7 @@ public unsafe Operation[] GetControlInputs() | |||
| 50 | 50 | { | |
| 51 | 51 | var control_inputs = new Operation[NumControlInputs]; | |
| 52 | 52 | ||
| 53 | - if(NumControlInputs > 0) | ||
| 53 | + if (NumControlInputs > 0) | ||
| 54 | 54 | { | |
| 55 | 55 | IntPtr control_input_handle = Marshal.AllocHGlobal(Marshal.SizeOf<IntPtr>() * NumControlInputs); | |
| 56 | 56 | c_api.TF_OperationGetControlInputs(_handle, control_input_handle, NumControlInputs); | |
@@ -70,7 +70,7 @@ public unsafe Operation[] GetControlOutputs() | |||
| 70 | 70 | { | |
| 71 | 71 | var control_outputs = new Operation[NumControlOutputs]; | |
| 72 | 72 | ||
| 73 | - if(NumControlOutputs > 0) | ||
| 73 | + if (NumControlOutputs > 0) | ||
| 74 | 74 | { | |
| 75 | 75 | IntPtr control_output_handle = Marshal.AllocHGlobal(Marshal.SizeOf<IntPtr>() * NumControlOutputs); | |
| 76 | 76 | c_api.TF_OperationGetControlOutputs(_handle, control_output_handle, NumControlInputs); | |
@@ -89,7 +89,7 @@ public Tensor[] outputs | |||
| 89 | 89 | { | |
| 90 | 90 | get | |
| 91 | 91 | { | |
| 92 | - if(_outputs == null) | ||
| 92 | + if (_outputs == null) | ||
| 93 | 93 | { | |
| 94 | 94 | _outputs = new Tensor[NumOutputs]; | |
| 95 | 95 | ||
@@ -106,7 +106,7 @@ public InputList inputs | |||
| 106 | 106 | { | |
| 107 | 107 | get | |
| 108 | 108 | { | |
| 109 | - if(_inputs == null) | ||
| 109 | + if (_inputs == null) | ||
| 110 | 110 | { | |
| 111 | 111 | var retval = new Tensor[NumInputs]; | |
| 112 | 112 | ||
@@ -124,6 +124,18 @@ public InputList inputs | |||
| 124 | 124 | } | |
| 125 | 125 | } | |
| 126 | 126 | ||
| 127 | + private NodeDef _node_def; | ||
| 128 | + public NodeDef node_def | ||
| 129 | + { | ||
| 130 | + get | ||
| 131 | + { | ||
| 132 | + if(_node_def == null) | ||
| 133 | + _node_def = GetNodeDef(); | ||
| 134 | + | ||
| 135 | + return _node_def; | ||
| 136 | + } | ||
| 137 | + } | ||
| 138 | + | ||
| 127 | 139 | public Operation(IntPtr handle) | |
| 128 | 140 | { | |
| 129 | 141 | if (handle == IntPtr.Zero) | |
@@ -195,7 +207,7 @@ public TF_AttrMetadata GetAttributeMetadata(string attr_name, Status s) | |||
| 195 | 207 | return c_api.TF_OperationGetAttrMetadata(_handle, attr_name, s); | |
| 196 | 208 | } | |
| 197 | 209 | ||
| 198 | - public NodeDef GetNodeDef() | ||
| 210 | + private NodeDef GetNodeDef() | ||
| 199 | 211 | { | |
| 200 | 212 | using (var s = new Status()) | |
| 201 | 213 | using (var buffer = new Buffer()) | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -36,6 +36,8 @@ public enum TF_DataType | |||
| 36 | 36 | TF_UINT32 = 22, | |
| 37 | 37 | TF_UINT64 = 23, | |
| 38 | 38 | ||
| 39 | + DtFloatRef = 101, // DT_FLOAT_REF | ||
| 39 | 40 | DtDoubleRef = 102, // DT_DOUBLE_REF | |
| 41 | + DtInt32Ref = 103, // DT_INT32_REF | ||
| 40 | 42 | } | |
| 41 | 43 | } | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -162,6 +162,7 @@ public Tensor(Operation op, int value_index, TF_DataType dtype) | |||
| 162 | 162 | this.op = op; | |
| 163 | 163 | this.value_index = value_index; | |
| 164 | 164 | this._dtype = dtype; | |
| 165 | + _id = ops.uid(); | ||
| 165 | 166 | } | |
| 166 | 167 | ||
| 167 | 168 | public List<Operation> consumers() | |
| 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 | |
@@ -69,14 +70,15 @@ private void _init_from_args(object initial_value, | |||
| 69 | 70 | { | |
| 70 | 71 | ||
| 71 | 72 | } | |
| 73 | + // Or get the initial value from a Tensor or Python object. | ||
| 72 | 74 | else | |
| 73 | 75 | { | |
| 74 | 76 | _initial_value = ops.convert_to_tensor(initial_value, name: "initial_value"); | |
| 75 | - } | ||
| 76 | 77 | ||
| 77 | - var shape = _initial_value.shape; | ||
| 78 | - dtype = _initial_value.dtype; | ||
| 79 | - _variable = gen_state_ops.variable_v2(shape, dtype, name); | ||
| 78 | + var shape = _initial_value.shape; | ||
| 79 | + dtype = _initial_value.dtype; | ||
| 80 | + _variable = gen_state_ops.variable_v2(shape, dtype, name); | ||
| 81 | + } | ||
| 80 | 82 | ||
| 81 | 83 | // Manually overrides the variable's shape with the initial value's. | |
| 82 | 84 | if (validate_shape) | |
@@ -87,8 +89,9 @@ private void _init_from_args(object initial_value, | |||
| 87 | 89 | // If 'initial_value' makes use of other variables, make sure we don't | |
| 88 | 90 | // have an issue if these other variables aren't initialized first by | |
| 89 | 91 | // using their initialized_value() method. | |
| 92 | + var _initial_value2 = _try_guard_against_uninitialized_dependencies(_initial_value); | ||
| 90 | 93 | ||
| 91 | - _initializer_op = gen_state_ops.assign(_variable, _initial_value, validate_shape).op; | ||
| 94 | + _initializer_op = gen_state_ops.assign(_variable, _initial_value2, validate_shape).op; | ||
| 92 | 95 | ||
| 93 | 96 | if (!String.IsNullOrEmpty(caching_device)) | |
| 94 | 97 | { | |
@@ -112,5 +115,51 @@ public Tensor _AsTensor() | |||
| 112 | 115 | { | |
| 113 | 116 | return _snapshot; | |
| 114 | 117 | } | |
| 118 | + | ||
| 119 | + /// <summary> | ||
| 120 | + /// Attempt to guard against dependencies on uninitialized variables. | ||
| 121 | + /// </summary> | ||
| 122 | + /// <param name="initial_value"></param> | ||
| 123 | + private Tensor _try_guard_against_uninitialized_dependencies(Tensor initial_value) | ||
| 124 | + { | ||
| 125 | + return _safe_initial_value_from_tensor(initial_value, new Dictionary<string, Operation>()); | ||
| 126 | + } | ||
| 127 | + | ||
| 128 | + /// <summary> | ||
| 129 | + /// Replace dependencies on variables with their initialized values. | ||
| 130 | + /// </summary> | ||
| 131 | + /// <param name="tensor">A `Tensor`. The tensor to replace.</param> | ||
| 132 | + /// <param name="op_cache">A dict mapping operation names to `Operation`s.</param> | ||
| 133 | + /// <returns>A `Tensor` compatible with `tensor`.</returns> | ||
| 134 | + private Tensor _safe_initial_value_from_tensor(Tensor tensor, Dictionary<string, Operation> op_cache) | ||
| 135 | + { | ||
| 136 | + var op = tensor.op; | ||
| 137 | + var new_op = op_cache.ContainsKey(op.Name) ? op_cache[op.Name] : null; | ||
| 138 | + if(new_op == null) | ||
| 139 | + { | ||
| 140 | + new_op = _safe_initial_value_from_op(op, op_cache); | ||
| 141 | + op_cache[op.Name] = new_op; | ||
| 142 | + } | ||
| 143 | + return new_op.outputs[tensor.value_index]; | ||
| 144 | + } | ||
| 145 | + | ||
| 146 | + private Operation _safe_initial_value_from_op(Operation op, Dictionary<string, Operation> op_cache) | ||
| 147 | + { | ||
| 148 | + var op_type = op.node_def.Op; | ||
| 149 | + switch (op_type) | ||
| 150 | + { | ||
| 151 | + case "IsVariableInitialized": | ||
| 152 | + case "VarIsInitializedOp": | ||
| 153 | + case "ReadVariableOp": | ||
| 154 | + return op; | ||
| 155 | + case "Variable": | ||
| 156 | + case "VariableV2": | ||
| 157 | + case "VarHandleOp": | ||
| 158 | + break; | ||
| 159 | + } | ||
| 160 | + | ||
| 161 | + // Recursively build initializer expressions for inputs. | ||
| 162 | + return op; | ||
| 163 | + } | ||
| 115 | 164 | } | |
| 116 | 165 | } | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -42,7 +42,7 @@ public static Tensor variable_v2(long[] shape, TF_DataType dtype, string name = | |||
| 42 | 42 | ||
| 43 | 43 | _execute.record_gradient("VariableV2", _inputs_flat, _attrs, _result, name); | |
| 44 | 44 | ||
| 45 | - return new Tensor(_op, 0, dtype); | ||
| 45 | + return _result[0]; | ||
| 46 | 46 | } | |
| 47 | 47 | ||
| 48 | 48 | /// <summary> | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -130,7 +130,7 @@ public void Graph() | |||
| 130 | 130 | EXPECT_EQ(TF_Code.TF_OK, s.Code); | |
| 131 | 131 | ||
| 132 | 132 | // Serialize to NodeDef. | |
| 133 | - var node_def = neg.GetNodeDef(); | ||
| 133 | + var node_def = neg.node_def; | ||
| 134 | 134 | ||
| 135 | 135 | // Validate NodeDef is what we expect. | |
| 136 | 136 | ASSERT_TRUE(c_test_util.IsNeg(node_def, "add")); | |
@@ -145,13 +145,13 @@ public void Graph() | |||
| 145 | 145 | // Look up some nodes by name. | |
| 146 | 146 | Operation neg2 = c_api.TF_GraphOperationByName(graph, "neg"); | |
| 147 | 147 | EXPECT_EQ(neg, neg2); | |
| 148 | - var node_def2 = neg2.GetNodeDef(); | ||
| 148 | + var node_def2 = neg2.node_def; | ||
| 149 | 149 | EXPECT_EQ(node_def.ToString(), node_def2.ToString()); | |
| 150 | 150 | ||
| 151 | 151 | Operation feed2 = c_api.TF_GraphOperationByName(graph, "feed"); | |
| 152 | 152 | EXPECT_EQ(feed, feed2); | |
| 153 | - node_def = feed.GetNodeDef(); | ||
| 154 | - node_def2 = feed2.GetNodeDef(); | ||
| 153 | + node_def = feed.node_def; | ||
| 154 | + node_def2 = feed2.node_def; | ||
| 155 | 155 | EXPECT_EQ(node_def.ToString(), node_def2.ToString()); | |
| 156 | 156 | ||
| 157 | 157 | // Test iterating through the nodes of a graph. | |
@@ -186,7 +186,7 @@ public void Graph() | |||
| 186 | 186 | } | |
| 187 | 187 | else | |
| 188 | 188 | { | |
| 189 | - node_def = oper.GetNodeDef(); | ||
| 189 | + node_def = oper.node_def; | ||
| 190 | 190 | Assert.Fail($"Unexpected Node: {node_def.ToString()}"); | |
| 191 | 191 | } | |
| 192 | 192 | } | |
| Back | FazBrowse Home | New Git URL |
0 commit comments