| FazBrowse GitHub Viewer | Trending | | Home |
| Tools: [Download Repo ZIP] [Original HTTPS Page] |
1 parent b02f285 commit e92aa44
23 files changed
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -1,4 +1,5 @@ | |||
| 1 | - using System; | ||
| 1 | + using NumSharp; | ||
| 2 | + using System; | ||
| 2 | 3 | using System.Collections.Generic; | |
| 3 | 4 | using System.Text; | |
| 4 | 5 | ||
@@ -7,5 +8,6 @@ namespace Tensorflow.Keras.ArgsDefinition | |||
| 7 | 8 | public class TensorFlowOpLayerArgs : LayerArgs | |
| 8 | 9 | { | |
| 9 | 10 | public NodeDef NodeDef { get; set; } | |
| 11 | + public Dictionary<int, NDArray> Constants { get; set; } | ||
| 10 | 12 | } | |
| 11 | 13 | } | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -160,9 +160,9 @@ public Tensor spatial_2d_padding(Tensor x, NDArray padding = null, string data_f | |||
| 160 | 160 | /// </summary> | |
| 161 | 161 | /// <param name="outputs"></param> | |
| 162 | 162 | /// <returns></returns> | |
| 163 | - public Tensor eval_in_eager_or_function(Tensor outputs) | ||
| 163 | + public NDArray eval_in_eager_or_function(Tensor outputs) | ||
| 164 | 164 | { | |
| 165 | - throw new NotImplementedException(""); | ||
| 165 | + return outputs.eval(); | ||
| 166 | 166 | } | |
| 167 | 167 | ||
| 168 | 168 | public class _DummyEagerGraph | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -21,7 +21,7 @@ public Flatten(FlattenArgs args) | |||
| 21 | 21 | _channels_first = args.DataFormat == "channels_first"; | |
| 22 | 22 | } | |
| 23 | 23 | ||
| 24 | - protected override Tensors call(Tensors inputs, Tensor state = null, bool is_training = false) | ||
| 24 | + protected override Tensors call_fn(Tensors inputs, Tensor state = null, bool is_training = false) | ||
| 25 | 25 | { | |
| 26 | 26 | if (_channels_first) | |
| 27 | 27 | { | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -69,10 +69,14 @@ void _init_graph_network(Tensors inputs, Tensors outputs) | |||
| 69 | 69 | } | |
| 70 | 70 | } | |
| 71 | 71 | ||
| 72 | - protected override Tensors call(Tensors inputs, Tensor state = null, bool is_training = false) | ||
| 72 | + protected override Tensors call_fn(Tensors inputs, Tensor state = null, bool is_training = false) | ||
| 73 | 73 | { | |
| 74 | - return base.call(inputs, state, is_training); | ||
| 74 | + return run_internal_graph(inputs, state, is_training); | ||
| 75 | 75 | } | |
| 76 | 76 | ||
| 77 | + Tensors run_internal_graph(Tensors inputs, Tensor state = null, bool is_training = false) | ||
| 78 | + { | ||
| 79 | + throw new NotImplementedException(""); | ||
| 80 | + } | ||
| 77 | 81 | } | |
| 78 | 82 | } | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -46,7 +46,7 @@ public Tensors Apply(Tensors inputs, Tensor state = null, bool is_training = fal | |||
| 46 | 46 | if (!built) | |
| 47 | 47 | MaybeBuild(inputs); | |
| 48 | 48 | ||
| 49 | - outputs = call(inputs, state: state, is_training: is_training); | ||
| 49 | + outputs = call_fn(inputs, state: state, is_training: is_training); | ||
| 50 | 50 | ||
| 51 | 51 | outputs = _set_connectivity_metadata_(inputs, outputs); | |
| 52 | 52 | _handle_activity_regularization(inputs, outputs); | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -42,7 +42,7 @@ Tensors FunctionalConstructionCall(Tensors inputs) | |||
| 42 | 42 | if (!dynamic) | |
| 43 | 43 | throw new NotImplementedException(""); | |
| 44 | 44 | ||
| 45 | - outputs = call(inputs); | ||
| 45 | + outputs = call_fn(inputs); | ||
| 46 | 46 | ||
| 47 | 47 | outputs = _set_connectivity_metadata_(inputs, outputs); | |
| 48 | 48 | _handle_activity_regularization(inputs, outputs); | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -162,7 +162,7 @@ private Tensor compute_mask(Tensor inputs, Tensor mask = null) | |||
| 162 | 162 | /// <param name="state"></param> | |
| 163 | 163 | /// <param name="is_training"></param> | |
| 164 | 164 | /// <returns></returns> | |
| 165 | - protected virtual Tensors call(Tensors inputs, Tensor state = null, bool is_training = false) | ||
| 165 | + protected virtual Tensors call_fn(Tensors inputs, Tensor state = null, bool is_training = false) | ||
| 166 | 166 | { | |
| 167 | 167 | throw new NotImplementedException(""); | |
| 168 | 168 | } | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -52,8 +52,6 @@ public Node(Layer layer, NodeArgs args) | |||
| 52 | 52 | layer.InboundNodes.Add(this); | |
| 53 | 53 | foreach (var kt in kerasInputs) | |
| 54 | 54 | { | |
| 55 | - if (kt.KerasHistory == null) | ||
| 56 | - continue; | ||
| 57 | 55 | var inbound_layer = kt.KerasHistory.layer; | |
| 58 | 56 | if (inbound_layer != null) | |
| 59 | 57 | inbound_layer.OutboundNodes.Add(this); | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -23,9 +23,9 @@ public TensorFlowOpLayer(TensorFlowOpLayerArgs args) | |||
| 23 | 23 | built = true; | |
| 24 | 24 | } | |
| 25 | 25 | ||
| 26 | - protected override Tensors call(Tensors inputs, Tensor state = null, bool is_training = false) | ||
| 26 | + protected override Tensors call_fn(Tensors inputs, Tensor state = null, bool is_training = false) | ||
| 27 | 27 | { | |
| 28 | - return base.call(inputs, state, is_training); | ||
| 28 | + return base.call_fn(inputs, state, is_training); | ||
| 29 | 29 | } | |
| 30 | 30 | } | |
| 31 | 31 | } | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -119,7 +119,7 @@ protected override void build(TensorShape input_shape) | |||
| 119 | 119 | built = true; | |
| 120 | 120 | } | |
| 121 | 121 | ||
| 122 | - protected override Tensors call(Tensors inputs, Tensor state = null, bool is_training = false) | ||
| 122 | + protected override Tensors call_fn(Tensors inputs, Tensor state = null, bool is_training = false) | ||
| 123 | 123 | { | |
| 124 | 124 | Tensor outputs = null; | |
| 125 | 125 | ||
| Back | FazBrowse Home | New Git URL |
0 commit comments