| FazBrowse GitHub Viewer | Trending | | Home |
| Tools: [Download Repo ZIP] [Original HTTPS Page] |
1 parent 0f7bf4d commit 321ddfc
15 files changed
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -12,20 +12,16 @@ public class SimpleRnnTest | |||
| 12 | 12 | { | |
| 13 | 13 | public void Run() | |
| 14 | 14 | { | |
| 15 | - tf.keras = new KerasInterface(); | ||
| 16 | - var inputs = np.random.random((32, 10, 8)).astype(np.float32); | ||
| 17 | - var simple_rnn = tf.keras.layers.SimpleRNN(4); | ||
| 18 | - var output = simple_rnn.Apply(inputs); // The output has shape `[32, 4]`. | ||
| 19 | - if (output.shape == (32, 4)) | ||
| 20 | - { | ||
| 15 | + tf.UseKeras<KerasInterface>(); | ||
| 16 | + var inputs = np.random.random((6, 10, 8)).astype(np.float32); | ||
| 17 | + //var simple_rnn = tf.keras.layers.SimpleRNN(4); | ||
| 18 | + //var output = simple_rnn.Apply(inputs); // The output has shape `[32, 4]`. | ||
| 21 | 19 | ||
| 22 | - } | ||
| 23 | - /*simple_rnn = tf.keras.layers.SimpleRNN( | ||
| 24 | - 4, return_sequences = True, return_state = True) | ||
| 20 | + var simple_rnn = tf.keras.layers.SimpleRNN(4, return_sequences: true, return_state: true); | ||
| 25 | 21 | ||
| 26 | - # whole_sequence_output has shape `[32, 10, 4]`. | ||
| 27 | - # final_state has shape `[32, 4]`. | ||
| 28 | - whole_sequence_output, final_state = simple_rnn(inputs)*/ | ||
| 22 | + // whole_sequence_output has shape `[32, 10, 4]`. | ||
| 23 | + // final_state has shape `[32, 4]`. | ||
| 24 | + var (whole_sequence_output, final_state) = simple_rnn.Apply(inputs); | ||
| 29 | 25 | } | |
| 30 | 26 | } | |
| 31 | 27 | } | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -9,6 +9,7 @@ public interface ILayer | |||
| 9 | 9 | string Name { get; } | |
| 10 | 10 | bool Trainable { get; } | |
| 11 | 11 | bool Built { get; } | |
| 12 | + void build(Shape input_shape); | ||
| 12 | 13 | List<ILayer> Layers { get; } | |
| 13 | 14 | List<INode> InboundNodes { get; } | |
| 14 | 15 | List<INode> OutboundNodes { get; } | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -163,7 +163,9 @@ public ILayer SimpleRNN(int units, | |||
| 163 | 163 | string activation = "tanh", | |
| 164 | 164 | string kernel_initializer = "glorot_uniform", | |
| 165 | 165 | string recurrent_initializer = "orthogonal", | |
| 166 | - string bias_initializer = "zeros"); | ||
| 166 | + string bias_initializer = "zeros", | ||
| 167 | + bool return_sequences = false, | ||
| 168 | + bool return_state = false); | ||
| 167 | 169 | ||
| 168 | 170 | public ILayer Subtract(); | |
| 169 | 171 | } | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -1,12 +1,32 @@ | |||
| 1 | 1 | using System; | |
| 2 | + using System.Linq; | ||
| 3 | + using static Tensorflow.TensorShapeProto.Types; | ||
| 2 | 4 | ||
| 3 | 5 | namespace Tensorflow.Operations.Initializers | |
| 4 | 6 | { | |
| 5 | 7 | public class Orthogonal : IInitializer | |
| 6 | 8 | { | |
| 9 | + float _gain = 0f; | ||
| 10 | + | ||
| 11 | + public Orthogonal(float gain = 1.0f, int? seed = null) | ||
| 12 | + { | ||
| 13 | + | ||
| 14 | + } | ||
| 15 | + | ||
| 7 | 16 | public Tensor Apply(InitializerArgs args) | |
| 8 | 17 | { | |
| 9 | - throw new NotImplementedException(); | ||
| 18 | + return _generate_init_val(args.Shape, args.DType); | ||
| 19 | + } | ||
| 20 | + | ||
| 21 | + private Tensor _generate_init_val(Shape shape, TF_DataType dtype) | ||
| 22 | + { | ||
| 23 | + var num_rows = 1L; | ||
| 24 | + foreach (var dim in shape.dims.Take(shape.ndim - 1)) | ||
| 25 | + num_rows *= dim; | ||
| 26 | + var num_cols = shape.dims.Last(); | ||
| 27 | + var flat_shape = (Math.Max(num_cols, num_rows), Math.Min(num_cols, num_rows)); | ||
| 28 | + | ||
| 29 | + throw new NotImplementedException(""); | ||
| 10 | 30 | } | |
| 11 | 31 | } | |
| 12 | 32 | } | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -147,5 +147,10 @@ public LayerArgs get_config() | |||
| 147 | 147 | { | |
| 148 | 148 | throw new NotImplementedException(); | |
| 149 | 149 | } | |
| 150 | + | ||
| 151 | + public void build(Shape input_shape) | ||
| 152 | + { | ||
| 153 | + throw new NotImplementedException(); | ||
| 154 | + } | ||
| 150 | 155 | } | |
| 151 | 156 | } | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -65,6 +65,11 @@ public tensorflow() | |||
| 65 | 65 | InitGradientEnvironment(); | |
| 66 | 66 | } | |
| 67 | 67 | ||
| 68 | + public void UseKeras<T>() where T : IKerasApi, new() | ||
| 69 | + { | ||
| 70 | + keras = new T(); | ||
| 71 | + } | ||
| 72 | + | ||
| 68 | 73 | public string VERSION => c_api.StringPiece(c_api.TF_Version()); | |
| 69 | 74 | ||
| 70 | 75 | private void InitGradientEnvironment() | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -65,7 +65,12 @@ protected void _init_graph_network(Tensors inputs, Tensors outputs) | |||
| 65 | 65 | } | |
| 66 | 66 | ||
| 67 | 67 | // Keep track of the network's nodes and layers. | |
| 68 | - (NetworkNodes, NodesByDepth, _self_tracked_trackables, _) = MapGraphNetwork(inputs, outputs); | ||
| 68 | + (NetworkNodes, NodesByDepth, var layers, _) = MapGraphNetwork(inputs, outputs); | ||
| 69 | + | ||
| 70 | + if (!_self_tracked_trackables.Any()) | ||
| 71 | + { | ||
| 72 | + _self_tracked_trackables = layers; | ||
| 73 | + } | ||
| 69 | 74 | ||
| 70 | 75 | // Build self.input_names and self.output_names. | |
| 71 | 76 | _set_output_names(); | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -1,9 +1,6 @@ | |||
| 1 | 1 | using System; | |
| 2 | 2 | using System.Linq; | |
| 3 | 3 | using Tensorflow.Graphs; | |
| 4 | - using Tensorflow.Keras.ArgsDefinition; | ||
| 5 | - using Tensorflow.Keras.Losses; | ||
| 6 | - using Tensorflow.Keras.Optimizers; | ||
| 7 | 4 | using static Tensorflow.Binding; | |
| 8 | 5 | using static Tensorflow.KerasApi; | |
| 9 | 6 | ||
@@ -13,6 +10,12 @@ public partial class Model | |||
| 13 | 10 | { | |
| 14 | 11 | public override void build(Shape input_shape) | |
| 15 | 12 | { | |
| 13 | + if (this is Functional || this is Sequential) | ||
| 14 | + { | ||
| 15 | + base.build(input_shape); | ||
| 16 | + return; | ||
| 17 | + } | ||
| 18 | + | ||
| 16 | 19 | var graph = tf.executing_eagerly() ? new FuncGraph("build_graph") : keras.backend.get_graph(); | |
| 17 | 20 | ||
| 18 | 21 | graph.as_default(); | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -122,15 +122,9 @@ public void add(ILayer layer) | |||
| 122 | 122 | else | |
| 123 | 123 | { | |
| 124 | 124 | _self_tracked_trackables.add(layer); | |
| 125 | - _handle_deferred_layer_dependencies(layer); | ||
| 126 | 125 | } | |
| 127 | 126 | } | |
| 128 | 127 | ||
| 129 | - void _handle_deferred_layer_dependencies(params ILayer[] layers) | ||
| 130 | - { | ||
| 131 | - _self_tracked_trackables.AddRange(layers); | ||
| 132 | - } | ||
| 133 | - | ||
| 134 | 128 | protected override Tensors Call(Tensors inputs, Tensor state = null, bool? training = null) | |
| 135 | 129 | { | |
| 136 | 130 | if (!_has_explicit_input_shape) | |
@@ -156,7 +150,7 @@ void _build_graph_network_for_inferred_shape(Shape input_shape, TF_DataType inpu | |||
| 156 | 150 | ops.init_scope(); | |
| 157 | 151 | var inputs = keras.Input(batch_input_shape: input_shape, | |
| 158 | 152 | dtype: input_dtype, | |
| 159 | - name: $"{_self_tracked_trackables[0].Name}_input"); | ||
| 153 | + name: _self_tracked_trackables[0].Name.EndsWith("_input") ? _self_tracked_trackables[0].Name : $"{_self_tracked_trackables[0].Name}_input"); | ||
| 160 | 154 | Tensors layer_input = inputs; | |
| 161 | 155 | Tensors layer_output = null; | |
| 162 | 156 | Tensors outputs = null; | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -658,14 +658,18 @@ public ILayer SimpleRNN(int units, | |||
| 658 | 658 | string activation = "tanh", | |
| 659 | 659 | string kernel_initializer = "glorot_uniform", | |
| 660 | 660 | string recurrent_initializer = "orthogonal", | |
| 661 | - string bias_initializer = "zeros") | ||
| 661 | + string bias_initializer = "zeros", | ||
| 662 | + bool return_sequences = false, | ||
| 663 | + bool return_state = false) | ||
| 662 | 664 | => new SimpleRNN(new SimpleRNNArgs | |
| 663 | 665 | { | |
| 664 | 666 | Units = units, | |
| 665 | 667 | Activation = GetActivationByName(activation), | |
| 666 | 668 | KernelInitializer = GetInitializerByName(kernel_initializer), | |
| 667 | - RecurrentInitializer= GetInitializerByName(recurrent_initializer), | ||
| 668 | - BiasInitializer= GetInitializerByName(bias_initializer) | ||
| 669 | + RecurrentInitializer = GetInitializerByName(recurrent_initializer), | ||
| 670 | + BiasInitializer = GetInitializerByName(bias_initializer), | ||
| 671 | + ReturnSequences = return_sequences, | ||
| 672 | + ReturnState = return_state | ||
| 669 | 673 | }); | |
| 670 | 674 | ||
| 671 | 675 | /// <summary> | |
| Back | FazBrowse Home | New Git URL |
0 commit comments