| FazBrowse GitHub Viewer | Trending | | Home |
| Tools: [Download Repo ZIP] [Original HTTPS Page] |
1 parent 090dc1e commit 93a242c
4 files changed
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -1,13 +1,15 @@ | |||
| 1 | - using System; | ||
| 1 | + using Newtonsoft.Json; | ||
| 2 | + using System; | ||
| 2 | 3 | using System.Collections.Generic; | |
| 3 | 4 | using System.Text; | |
| 4 | 5 | ||
| 5 | 6 | namespace Tensorflow.Keras.ArgsDefinition | |
| 6 | 7 | { | |
| 7 | 8 | // TODO: complete the implementation | |
| 8 | - public class MergeArgs : LayerArgs | ||
| 9 | + public class MergeArgs : AutoSerializeLayerArgs | ||
| 9 | 10 | { | |
| 10 | 11 | public Tensors Inputs { get; set; } | |
| 12 | + [JsonProperty("axis")] | ||
| 11 | 13 | public int Axis { get; set; } | |
| 12 | 14 | } | |
| 13 | 15 | } | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -30,7 +30,7 @@ public static (Tensors, Tensors, Dictionary<string, ILayer>) reconstruct_from_co | |||
| 30 | 30 | created_layers = created_layers ?? new Dictionary<string, ILayer>(); | |
| 31 | 31 | var node_index_map = new Dictionary<(string, int), int>(); | |
| 32 | 32 | var node_count_by_layer = new Dictionary<ILayer, int>(); | |
| 33 | - var unprocessed_nodes = new Dictionary<ILayer, NodeConfig>(); | ||
| 33 | + var unprocessed_nodes = new Dictionary<ILayer, List<NodeConfig>>(); | ||
| 34 | 34 | // First, we create all layers and enqueue nodes to be processed | |
| 35 | 35 | foreach (var layer_data in config.Layers) | |
| 36 | 36 | process_layer(created_layers, layer_data, unprocessed_nodes, node_count_by_layer); | |
@@ -79,7 +79,7 @@ public static (Tensors, Tensors, Dictionary<string, ILayer>) reconstruct_from_co | |||
| 79 | 79 | ||
| 80 | 80 | static void process_layer(Dictionary<string, ILayer> created_layers, | |
| 81 | 81 | LayerConfig layer_data, | |
| 82 | - Dictionary<ILayer, NodeConfig> unprocessed_nodes, | ||
| 82 | + Dictionary<ILayer, List<NodeConfig>> unprocessed_nodes, | ||
| 83 | 83 | Dictionary<ILayer, int> node_count_by_layer) | |
| 84 | 84 | { | |
| 85 | 85 | ILayer layer = null; | |
@@ -92,32 +92,38 @@ static void process_layer(Dictionary<string, ILayer> created_layers, | |||
| 92 | 92 | ||
| 93 | 93 | created_layers[layer_name] = layer; | |
| 94 | 94 | } | |
| 95 | - node_count_by_layer[layer] = _should_skip_first_node(layer) ? 1 : 0; | ||
| 95 | + node_count_by_layer[layer] = layer_data.InboundNodes.Count - (_should_skip_first_node(layer) ? 1 : 0); | ||
| 96 | 96 | ||
| 97 | 97 | var inbound_nodes_data = layer_data.InboundNodes; | |
| 98 | 98 | foreach (var node_data in inbound_nodes_data) | |
| 99 | 99 | { | |
| 100 | 100 | if (!unprocessed_nodes.ContainsKey(layer)) | |
| 101 | - unprocessed_nodes[layer] = node_data; | ||
| 101 | + unprocessed_nodes[layer] = new List<NodeConfig>() { node_data }; | ||
| 102 | 102 | else | |
| 103 | - unprocessed_nodes.Add(layer, node_data); | ||
| 103 | + unprocessed_nodes[layer].Add(node_data); | ||
| 104 | 104 | } | |
| 105 | 105 | } | |
| 106 | 106 | ||
| 107 | 107 | static void process_node(ILayer layer, | |
| 108 | - NodeConfig node_data, | ||
| 108 | + List<NodeConfig> nodes_data, | ||
| 109 | 109 | Dictionary<string, ILayer> created_layers, | |
| 110 | 110 | Dictionary<ILayer, int> node_count_by_layer, | |
| 111 | 111 | Dictionary<(string, int), int> node_index_map) | |
| 112 | 112 | { | |
| 113 | + | ||
| 113 | 114 | var input_tensors = new List<Tensor>(); | |
| 114 | - var inbound_layer_name = node_data.Name; | ||
| 115 | - var inbound_node_index = node_data.NodeIndex; | ||
| 116 | - var inbound_tensor_index = node_data.TensorIndex; | ||
| 117 | 115 | ||
| 118 | - var inbound_layer = created_layers[inbound_layer_name]; | ||
| 119 | - var inbound_node = inbound_layer.InboundNodes[inbound_node_index]; | ||
| 120 | - input_tensors.Add(inbound_node.Outputs[inbound_node_index]); | ||
| 116 | + for (int i = 0; i < nodes_data.Count; i++) | ||
| 117 | + { | ||
| 118 | + var node_data = nodes_data[i]; | ||
| 119 | + var inbound_layer_name = node_data.Name; | ||
| 120 | + var inbound_node_index = node_data.NodeIndex; | ||
| 121 | + var inbound_tensor_index = node_data.TensorIndex; | ||
| 122 | + | ||
| 123 | + var inbound_layer = created_layers[inbound_layer_name]; | ||
| 124 | + var inbound_node = inbound_layer.InboundNodes[inbound_node_index]; | ||
| 125 | + input_tensors.Add(inbound_node.Outputs[inbound_node_index]); | ||
| 126 | + } | ||
| 121 | 127 | ||
| 122 | 128 | var output_tensors = layer.Apply(input_tensors); | |
| 123 | 129 | ||
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -39,6 +39,7 @@ public override void build(KerasShapesWrapper input_shape) | |||
| 39 | 39 | shape_set.Add(shape); | |
| 40 | 40 | }*/ | |
| 41 | 41 | _buildInputShape = input_shape; | |
| 42 | + built = true; | ||
| 42 | 43 | } | |
| 43 | 44 | ||
| 44 | 45 | protected override Tensors _merge_function(Tensors inputs) | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -112,12 +112,23 @@ public static FunctionalConfig deserialize_model_config(JToken json) | |||
| 112 | 112 | foreach (var token in layersToken) | |
| 113 | 113 | { | |
| 114 | 114 | var args = deserialize_layer_args(token["class_name"].ToObject<string>(), token["config"]); | |
| 115 | + | ||
| 116 | + List<NodeConfig> nodeConfig = null; //python tensorflow sometimes exports inbound nodes in an extra nested array | ||
| 117 | + if (token["inbound_nodes"].Count() > 0 && token["inbound_nodes"][0].Count() > 0 && token["inbound_nodes"][0][0].Count() > 0) | ||
| 118 | + { | ||
| 119 | + nodeConfig = token["inbound_nodes"].ToObject<List<List<NodeConfig>>>().FirstOrDefault() ?? new List<NodeConfig>(); | ||
| 120 | + } | ||
| 121 | + else | ||
| 122 | + { | ||
| 123 | + nodeConfig = token["inbound_nodes"].ToObject<List<NodeConfig>>(); | ||
| 124 | + } | ||
| 125 | + | ||
| 115 | 126 | config.Layers.Add(new LayerConfig() | |
| 116 | 127 | { | |
| 117 | 128 | Config = args, | |
| 118 | 129 | Name = token["name"].ToObject<string>(), | |
| 119 | 130 | ClassName = token["class_name"].ToObject<string>(), | |
| 120 | - InboundNodes = token["inbound_nodes"].ToObject<List<NodeConfig>>() | ||
| 131 | + InboundNodes = nodeConfig, | ||
| 121 | 132 | }); | |
| 122 | 133 | } | |
| 123 | 134 | config.InputLayers = json["input_layers"].ToObject<List<NodeConfig>>(); | |
| Back | FazBrowse Home | New Git URL |
0 commit comments