| FazBrowse GitHub Viewer | Trending | | Home |
| Tools: [Download Repo ZIP] [Original HTTPS Page] |
1 parent 05dd652 commit 58d2dae
4 files changed
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -17,8 +17,8 @@ public class FuncGraph : Graph | |||
| 17 | 17 | IntPtr func_handle; | |
| 18 | 18 | public string FuncName => _graph_key; | |
| 19 | 19 | ||
| 20 | - public Tensors Inputs { get; set; } | ||
| 21 | - public Tensors Outputs { get; set; } | ||
| 20 | + public Tensors Inputs { get; set; } = new Tensors(); | ||
| 21 | + public Tensors Outputs { get; set; } = new Tensors(); | ||
| 22 | 22 | public Dictionary<string, string> Attrs { get; set; } | |
| 23 | 23 | ||
| 24 | 24 | public Dictionary<long, (Tensor, Tensor)> _captures | |
@@ -175,14 +175,7 @@ Tensor _capture_helper(Tensor tensor, string name, TensorShape shape = null) | |||
| 175 | 175 | void add_capture(Tensor tensor, Tensor placeholder) | |
| 176 | 176 | { | |
| 177 | 177 | _captures.Add(tensor.Id, (tensor, placeholder)); | |
| 178 | - if (Inputs == null) | ||
| 179 | - Inputs = new Tensors(placeholder); | ||
| 180 | - else | ||
| 181 | - { | ||
| 182 | - var inputs = Inputs.ToList(); | ||
| 183 | - inputs.Add(placeholder); | ||
| 184 | - Inputs = new Tensors(inputs.ToArray()); | ||
| 185 | - } | ||
| 178 | + Inputs.Add(placeholder); | ||
| 186 | 179 | } | |
| 187 | 180 | ||
| 188 | 181 | Tensor _create_substitute_placeholder(Tensor value, | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -39,7 +39,8 @@ public class BaseSession : DisposableObject | |||
| 39 | 39 | public BaseSession(string target = "", Graph g = null, ConfigProto config = null, Status status = null) | |
| 40 | 40 | { | |
| 41 | 41 | _graph = g ?? ops.get_default_graph(); | |
| 42 | - _graph.as_default(); | ||
| 42 | + if (!_graph.building_function) | ||
| 43 | + _graph.as_default(); | ||
| 43 | 44 | _target = Encoding.UTF8.GetBytes(target); | |
| 44 | 45 | ||
| 45 | 46 | using (var opts = new SessionOptions(target, config)) | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -58,9 +58,6 @@ public InputLayer(InputLayerArgs args) : | |||
| 58 | 58 | args.DType = args.InputTensor == null ? tf.float32 : args.InputTensor.dtype; | |
| 59 | 59 | } | |
| 60 | 60 | ||
| 61 | - // In graph mode, create a graph placeholder to call the layer on. | ||
| 62 | - tf.Context.graph_mode(); | ||
| 63 | - | ||
| 64 | 61 | if (args.InputTensor == null) | |
| 65 | 62 | { | |
| 66 | 63 | if (args.InputShape != null) | |
@@ -74,15 +71,18 @@ public InputLayer(InputLayerArgs args) : | |||
| 74 | 71 | args.BatchInputShape = null; | |
| 75 | 72 | } | |
| 76 | 73 | ||
| 74 | + var graph = keras.backend.get_graph(); | ||
| 75 | + graph.as_default(); | ||
| 76 | + | ||
| 77 | 77 | args.InputTensor = keras.backend.placeholder( | |
| 78 | 78 | shape: BatchInputShape, | |
| 79 | 79 | dtype: DType, | |
| 80 | 80 | name: Name, | |
| 81 | 81 | sparse: args.Sparse, | |
| 82 | 82 | ragged: args.Ragged); | |
| 83 | 83 | ||
| 84 | - | ||
| 85 | 84 | isPlaceholder = true; | |
| 85 | + tf.Context.restore_mode(); | ||
| 86 | 86 | } | |
| 87 | 87 | ||
| 88 | 88 | // Create an input node to add to self.outbound_node | |
@@ -97,8 +97,6 @@ public InputLayer(InputLayerArgs args) : | |||
| 97 | 97 | typeSpec = new TensorSpec(args.InputTensor.TensorShape, | |
| 98 | 98 | dtype: args.InputTensor.dtype, | |
| 99 | 99 | name: Name); | |
| 100 | - | ||
| 101 | - tf.Context.restore_mode(); | ||
| 102 | 100 | } | |
| 103 | 101 | ||
| 104 | 102 | public static InputLayer from_config(LayerArgs args) | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -151,23 +151,12 @@ public static void CreateKerasHistoryHelper(Tensors tensors, List<Operation> pro | |||
| 151 | 151 | ||
| 152 | 152 | // recursively | |
| 153 | 153 | CreateKerasHistoryHelper(layer_inputs, processed_ops, created_layers); | |
| 154 | - Layer op_layer = null; | ||
| 155 | - /*var op_layer = new TensorFlowOpLayer(new TensorFlowOpLayerArgs | ||
| 154 | + Layer op_layer = new TensorFlowOpLayer(new TensorFlowOpLayerArgs | ||
| 156 | 155 | { | |
| 157 | 156 | NodeDef = op.node_def, | |
| 158 | 157 | Constants = constants, | |
| 159 | 158 | Name = op.name | |
| 160 | - });*/ | ||
| 161 | - op_layer = op.type switch | ||
| 162 | - { | ||
| 163 | - // "AddV2" => keras.layers.Add(), | ||
| 164 | - _ => new TensorFlowOpLayer(new TensorFlowOpLayerArgs | ||
| 165 | - { | ||
| 166 | - NodeDef = op.node_def, | ||
| 167 | - Constants = constants, | ||
| 168 | - Name = op.name | ||
| 169 | - }) | ||
| 170 | - }; | ||
| 159 | + }); | ||
| 171 | 160 | created_layers.Add(op_layer); | |
| 172 | 161 | op_layer.SetConnectivityMetadata(layer_inputs, op.outputs); | |
| 173 | 162 | processed_ops.Add(op); | |
| Back | FazBrowse Home | New Git URL |
0 commit comments