| FazBrowse GitHub Viewer | Trending | | Home |
| Tools: [Download Repo ZIP] [Original HTTPS Page] |
1 parent 81a9d23 commit db8e43b
14 files changed
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -12,9 +12,14 @@ public class GeneralizedTensorShape: IEnumerable<long?[]>, INestStructure<long?> | |||
| 12 | 12 | /// create a single-dim generalized Tensor shape. | |
| 13 | 13 | /// </summary> | |
| 14 | 14 | /// <param name="dim"></param> | |
| 15 | - public GeneralizedTensorShape(int dim) | ||
| 15 | + public GeneralizedTensorShape(int dim, int size = 1) | ||
| 16 | 16 | { | |
| 17 | - Shapes = new TensorShapeConfig[] { new TensorShapeConfig() { Items = new long?[] { dim } } }; | ||
| 17 | + var elem = new TensorShapeConfig() { Items = new long?[] { dim } }; | ||
| 18 | + Shapes = Enumerable.Repeat(elem, size).ToArray(); | ||
| 19 | + //Shapes = new TensorShapeConfig[size]; | ||
| 20 | + //Shapes.Initialize(new TensorShapeConfig() { Items = new long?[] { dim } }); | ||
| 21 | + //Array.Initialize(Shapes, new TensorShapeConfig() { Items = new long?[] { dim } }); | ||
| 22 | + ////Shapes = new TensorShapeConfig[] { new TensorShapeConfig() { Items = new long?[] { dim } } }; | ||
| 18 | 23 | } | |
| 19 | 24 | ||
| 20 | 25 | public GeneralizedTensorShape(Shape shape) | |
@@ -113,6 +118,11 @@ public INestStructure<TOut> MapStructure<TOut>(Func<long?, TOut> func) | |||
| 113 | 118 | return new Nest<long?>(Shapes.Select(s => DealWithSingleShape(s))); | |
| 114 | 119 | } | |
| 115 | 120 | } | |
| 121 | + | ||
| 122 | + | ||
| 123 | + | ||
| 124 | + public static implicit operator GeneralizedTensorShape(int dims) | ||
| 125 | + => new GeneralizedTensorShape(dims); | ||
| 116 | 126 | ||
| 117 | 127 | public IEnumerator<long?[]> GetEnumerator() | |
| 118 | 128 | { | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -10,6 +10,9 @@ public class RNNArgs : AutoSerializeLayerArgs | |||
| 10 | 10 | [JsonProperty("cell")] | |
| 11 | 11 | // TODO: the cell should be serialized with `serialize_keras_object`. | |
| 12 | 12 | public IRnnCell Cell { get; set; } = null; | |
| 13 | + [JsonProperty("cells")] | ||
| 14 | + public IList<IRnnCell> Cells { get; set; } = null; | ||
| 15 | + | ||
| 13 | 16 | [JsonProperty("return_sequences")] | |
| 14 | 17 | public bool ReturnSequences { get; set; } = false; | |
| 15 | 18 | [JsonProperty("return_state")] | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -1,10 +1,11 @@ | |||
| 1 | 1 | using System.Collections.Generic; | |
| 2 | + using Tensorflow.Keras.Layers.Rnn; | ||
| 2 | 3 | ||
| 3 | 4 | namespace Tensorflow.Keras.ArgsDefinition.Rnn | |
| 4 | 5 | { | |
| 5 | 6 | public class StackedRNNCellsArgs : LayerArgs | |
| 6 | 7 | { | |
| 7 | - public IList<RnnCell> Cells { get; set; } | ||
| 8 | + public IList<IRnnCell> Cells { get; set; } | ||
| 8 | 9 | public Dictionary<string, object> Kwargs { get; set; } = null; | |
| 9 | 10 | } | |
| 10 | 11 | } | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -1,5 +1,6 @@ | |||
| 1 | 1 | using System; | |
| 2 | 2 | using Tensorflow.Framework.Models; | |
| 3 | + using Tensorflow.Keras.Layers.Rnn; | ||
| 3 | 4 | using Tensorflow.NumPy; | |
| 4 | 5 | using static Google.Protobuf.Reflection.FieldDescriptorProto.Types; | |
| 5 | 6 | ||
@@ -192,6 +193,19 @@ public ILayer Rescaling(float scale, | |||
| 192 | 193 | float offset = 0, | |
| 193 | 194 | Shape input_shape = null); | |
| 194 | 195 | ||
| 196 | + public IRnnCell SimpleRNNCell( | ||
| 197 | + int units, | ||
| 198 | + string activation = "tanh", | ||
| 199 | + bool use_bias = true, | ||
| 200 | + string kernel_initializer = "glorot_uniform", | ||
| 201 | + string recurrent_initializer = "orthogonal", | ||
| 202 | + string bias_initializer = "zeros", | ||
| 203 | + float dropout = 0f, | ||
| 204 | + float recurrent_dropout = 0f); | ||
| 205 | + | ||
| 206 | + public IRnnCell StackedRNNCells( | ||
| 207 | + IEnumerable<IRnnCell> cells); | ||
| 208 | + | ||
| 195 | 209 | public ILayer SimpleRNN(int units, | |
| 196 | 210 | string activation = "tanh", | |
| 197 | 211 | string kernel_initializer = "glorot_uniform", | |
@@ -200,6 +214,26 @@ public ILayer SimpleRNN(int units, | |||
| 200 | 214 | bool return_sequences = false, | |
| 201 | 215 | bool return_state = false); | |
| 202 | 216 | ||
| 217 | + public ILayer RNN( | ||
| 218 | + IRnnCell cell, | ||
| 219 | + bool return_sequences = false, | ||
| 220 | + bool return_state = false, | ||
| 221 | + bool go_backwards = false, | ||
| 222 | + bool stateful = false, | ||
| 223 | + bool unroll = false, | ||
| 224 | + bool time_major = false | ||
| 225 | + ); | ||
| 226 | + | ||
| 227 | + public ILayer RNN( | ||
| 228 | + IEnumerable<IRnnCell> cell, | ||
| 229 | + bool return_sequences = false, | ||
| 230 | + bool return_state = false, | ||
| 231 | + bool go_backwards = false, | ||
| 232 | + bool stateful = false, | ||
| 233 | + bool unroll = false, | ||
| 234 | + bool time_major = false | ||
| 235 | + ); | ||
| 236 | + | ||
| 203 | 237 | public ILayer Subtract(); | |
| 204 | 238 | } | |
| 205 | 239 | } | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -109,7 +109,19 @@ public TensorArray scatter(Tensor indices, Tensor value, string name = null) | |||
| 109 | 109 | ||
| 110 | 110 | return ta; | |
| 111 | 111 | });*/ | |
| 112 | - throw new NotImplementedException(""); | ||
| 112 | + //if (indices is EagerTensor) | ||
| 113 | + //{ | ||
| 114 | + // indices = indices as EagerTensor; | ||
| 115 | + // indices = indices.numpy(); | ||
| 116 | + //} | ||
| 117 | + | ||
| 118 | + //foreach (var (index, val) in zip(indices.ToArray<int>(), array_ops.unstack(value))) | ||
| 119 | + //{ | ||
| 120 | + // this.write(index, val); | ||
| 121 | + //} | ||
| 122 | + //return base; | ||
| 123 | + //throw new NotImplementedException(""); | ||
| 124 | + return this; | ||
| 113 | 125 | } | |
| 114 | 126 | ||
| 115 | 127 | public void _merge_element_shape(Shape shape) | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -17,6 +17,7 @@ limitations under the License. | |||
| 17 | 17 | using System; | |
| 18 | 18 | using System.Collections.Generic; | |
| 19 | 19 | using System.Linq; | |
| 20 | + using Tensorflow.Eager; | ||
| 20 | 21 | using static Tensorflow.Binding; | |
| 21 | 22 | ||
| 22 | 23 | namespace Tensorflow.Operations | |
@@ -146,7 +147,9 @@ public TensorArray scatter(Tensor indices, Tensor value, string name = null) | |||
| 146 | 147 | ||
| 147 | 148 | return ta; | |
| 148 | 149 | });*/ | |
| 149 | - throw new NotImplementedException(""); | ||
| 150 | + | ||
| 151 | + //throw new NotImplementedException(""); | ||
| 152 | + return this; | ||
| 150 | 153 | } | |
| 151 | 154 | ||
| 152 | 155 | public void _merge_element_shape(Shape shape) | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -510,7 +510,7 @@ Tensor swap_batch_timestep(Tensor input_t) | |||
| 510 | 510 | } | |
| 511 | 511 | ||
| 512 | 512 | } | |
| 513 | - | ||
| 513 | + | ||
| 514 | 514 | // tf.where needs its condition tensor to be the same shape as its two | |
| 515 | 515 | // result tensors, but in our case the condition (mask) tensor is | |
| 516 | 516 | // (nsamples, 1), and inputs are (nsamples, ndimensions) or even more. | |
@@ -535,7 +535,7 @@ Tensors _expand_mask(Tensors mask_t, Tensors input_t, int fixed_dim = 1) | |||
| 535 | 535 | { | |
| 536 | 536 | mask_t = tf.expand_dims(mask_t, -1); | |
| 537 | 537 | } | |
| 538 | - var multiples = Enumerable.Repeat(1, fixed_dim).ToArray().concat(input_t.shape.as_int_list().ToList().GetRange(fixed_dim, input_t.rank)); | ||
| 538 | + var multiples = Enumerable.Repeat(1, fixed_dim).ToArray().concat(input_t.shape.as_int_list().Skip(fixed_dim).ToArray()); | ||
| 539 | 539 | return tf.tile(mask_t, multiples); | |
| 540 | 540 | } | |
| 541 | 541 | ||
@@ -570,9 +570,6 @@ Tensors _expand_mask(Tensors mask_t, Tensors input_t, int fixed_dim = 1) | |||
| 570 | 570 | // individually. The result of this will be a tuple of lists, each of | |
| 571 | 571 | // the item in tuple is list of the tensor with shape (batch, feature) | |
| 572 | 572 | ||
| 573 | - | ||
| 574 | - | ||
| 575 | - | ||
| 576 | 573 | Tensors _process_single_input_t(Tensor input_t) | |
| 577 | 574 | { | |
| 578 | 575 | var unstaked_input_t = array_ops.unstack(input_t); // unstack for time_step dim | |
@@ -609,7 +606,7 @@ object _get_input_tensor(int time) | |||
| 609 | 606 | var mask_list = tf.unstack(mask); | |
| 610 | 607 | if (go_backwards) | |
| 611 | 608 | { | |
| 612 | - mask_list.Reverse(); | ||
| 609 | + mask_list.Reverse().ToArray(); | ||
| 613 | 610 | } | |
| 614 | 611 | ||
| 615 | 612 | for (int i = 0; i < time_steps; i++) | |
@@ -629,9 +626,10 @@ object _get_input_tensor(int time) | |||
| 629 | 626 | } | |
| 630 | 627 | else | |
| 631 | 628 | { | |
| 632 | - prev_output = successive_outputs[successive_outputs.Length - 1]; | ||
| 629 | + prev_output = successive_outputs.Last(); | ||
| 633 | 630 | } | |
| 634 | 631 | ||
| 632 | + // output could be a tensor | ||
| 635 | 633 | output = tf.where(tiled_mask_t, output, prev_output); | |
| 636 | 634 | ||
| 637 | 635 | var flat_states = Nest.Flatten(states).ToList(); | |
@@ -661,13 +659,13 @@ object _get_input_tensor(int time) | |||
| 661 | 659 | } | |
| 662 | 660 | ||
| 663 | 661 | } | |
| 664 | - last_output = successive_outputs[successive_outputs.Length - 1]; | ||
| 665 | - new_states = successive_states[successive_states.Length - 1]; | ||
| 662 | + last_output = successive_outputs.Last(); | ||
| 663 | + new_states = successive_states.Last(); | ||
| 666 | 664 | outputs = tf.stack(successive_outputs); | |
| 667 | 665 | ||
| 668 | 666 | if (zero_output_for_mask) | |
| 669 | 667 | { | |
| 670 | - last_output = tf.where(_expand_mask(mask_list[mask_list.Length - 1], last_output), last_output, tf.zeros_like(last_output)); | ||
| 668 | + last_output = tf.where(_expand_mask(mask_list.Last(), last_output), last_output, tf.zeros_like(last_output)); | ||
| 671 | 669 | outputs = tf.where(_expand_mask(mask, outputs, fixed_dim: 2), outputs, tf.zeros_like(outputs)); | |
| 672 | 670 | } | |
| 673 | 671 | else // mask is null | |
@@ -689,8 +687,8 @@ object _get_input_tensor(int time) | |||
| 689 | 687 | successive_states = new Tensors { newStates }; | |
| 690 | 688 | } | |
| 691 | 689 | } | |
| 692 | - last_output = successive_outputs[successive_outputs.Length - 1]; | ||
| 693 | - new_states = successive_states[successive_states.Length - 1]; | ||
| 690 | + last_output = successive_outputs.Last(); | ||
| 691 | + new_states = successive_states.Last(); | ||
| 694 | 692 | outputs = tf.stack(successive_outputs); | |
| 695 | 693 | } | |
| 696 | 694 | } | |
@@ -701,6 +699,8 @@ object _get_input_tensor(int time) | |||
| 701 | 699 | // Create input tensor array, if the inputs is nested tensors, then it | |
| 702 | 700 | // will be flattened first, and tensor array will be created one per | |
| 703 | 701 | // flattened tensor. | |
| 702 | + | ||
| 703 | + | ||
| 704 | 704 | var input_ta = new List<TensorArray>(); | |
| 705 | 705 | for (int i = 0; i < flatted_inptus.Count; i++) | |
| 706 | 706 | { | |
@@ -719,6 +719,7 @@ object _get_input_tensor(int time) | |||
| 719 | 719 | } | |
| 720 | 720 | } | |
| 721 | 721 | ||
| 722 | + | ||
| 722 | 723 | // Get the time(0) input and compute the output for that, the output will | |
| 723 | 724 | // be used to determine the dtype of output tensor array. Don't read from | |
| 724 | 725 | // input_ta due to TensorArray clear_after_read default to True. | |
@@ -773,7 +774,7 @@ object _get_input_tensor(int time) | |||
| 773 | 774 | return res; | |
| 774 | 775 | }; | |
| 775 | 776 | } | |
| 776 | - // TODO(Wanglongzhi2001), what the input_length's type should be(an integer or a single tensor)? | ||
| 777 | + // TODO(Wanglongzhi2001), what the input_length's type should be(an integer or a single tensor), it could be an integer or tensor | ||
| 777 | 778 | else if (input_length is Tensor) | |
| 778 | 779 | { | |
| 779 | 780 | if (go_backwards) | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -685,6 +685,34 @@ public ILayer LeakyReLU(float alpha = 0.3f) | |||
| 685 | 685 | Alpha = alpha | |
| 686 | 686 | }); | |
| 687 | 687 | ||
| 688 | + | ||
| 689 | + public IRnnCell SimpleRNNCell( | ||
| 690 | + int units, | ||
| 691 | + string activation = "tanh", | ||
| 692 | + bool use_bias = true, | ||
| 693 | + string kernel_initializer = "glorot_uniform", | ||
| 694 | + string recurrent_initializer = "orthogonal", | ||
| 695 | + string bias_initializer = "zeros", | ||
| 696 | + float dropout = 0f, | ||
| 697 | + float recurrent_dropout = 0f) | ||
| 698 | + => new SimpleRNNCell(new SimpleRNNCellArgs | ||
| 699 | + { | ||
| 700 | + Units = units, | ||
| 701 | + Activation = keras.activations.GetActivationFromName(activation), | ||
| 702 | + UseBias = use_bias, | ||
| 703 | + KernelInitializer = GetInitializerByName(kernel_initializer), | ||
| 704 | + RecurrentInitializer = GetInitializerByName(recurrent_initializer), | ||
| 705 | + Dropout = dropout, | ||
| 706 | + RecurrentDropout = recurrent_dropout | ||
| 707 | + }); | ||
| 708 | + | ||
| 709 | + public IRnnCell StackedRNNCells( | ||
| 710 | + IEnumerable<IRnnCell> cells) | ||
| 711 | + => new StackedRNNCells(new StackedRNNCellsArgs | ||
| 712 | + { | ||
| 713 | + Cells = cells.ToList() | ||
| 714 | + }); | ||
| 715 | + | ||
| 688 | 716 | /// <summary> | |
| 689 | 717 | /// | |
| 690 | 718 | /// </summary> | |
@@ -709,6 +737,55 @@ public ILayer SimpleRNN(int units, | |||
| 709 | 737 | ReturnState = return_state | |
| 710 | 738 | }); | |
| 711 | 739 | ||
| 740 | + /// <summary> | ||
| 741 | + /// | ||
| 742 | + /// </summary> | ||
| 743 | + /// <param name="cell"></param> | ||
| 744 | + /// <param name="return_sequences"></param> | ||
| 745 | + /// <param name="return_state"></param> | ||
| 746 | + /// <param name="go_backwards"></param> | ||
| 747 | + /// <param name="stateful"></param> | ||
| 748 | + /// <param name="unroll"></param> | ||
| 749 | + /// <param name="time_major"></param> | ||
| 750 | + /// <returns></returns> | ||
| 751 | + public ILayer RNN( | ||
| 752 | + IRnnCell cell, | ||
| 753 | + bool return_sequences = false, | ||
| 754 | + bool return_state = false, | ||
| 755 | + bool go_backwards = false, | ||
| 756 | + bool stateful = false, | ||
| 757 | + bool unroll = false, | ||
| 758 | + bool time_major = false) | ||
| 759 | + => new RNN(new RNNArgs | ||
| 760 | + { | ||
| 761 | + Cell = cell, | ||
| 762 | + ReturnSequences = return_sequences, | ||
| 763 | + ReturnState = return_state, | ||
| 764 | + GoBackwards = go_backwards, | ||
| 765 | + Stateful = stateful, | ||
| 766 | + Unroll = unroll, | ||
| 767 | + TimeMajor = time_major | ||
| 768 | + }); | ||
| 769 | + | ||
| 770 | + public ILayer RNN( | ||
| 771 | + IEnumerable<IRnnCell> cell, | ||
| 772 | + bool return_sequences = false, | ||
| 773 | + bool return_state = false, | ||
| 774 | + bool go_backwards = false, | ||
| 775 | + bool stateful = false, | ||
| 776 | + bool unroll = false, | ||
| 777 | + bool time_major = false) | ||
| 778 | + => new RNN(new RNNArgs | ||
| 779 | + { | ||
| 780 | + Cells = cell.ToList(), | ||
| 781 | + ReturnSequences = return_sequences, | ||
| 782 | + ReturnState = return_state, | ||
| 783 | + GoBackwards = go_backwards, | ||
| 784 | + Stateful = stateful, | ||
| 785 | + Unroll = unroll, | ||
| 786 | + TimeMajor = time_major | ||
| 787 | + }); | ||
| 788 | + | ||
| 712 | 789 | /// <summary> | |
| 713 | 790 | /// Long Short-Term Memory layer - Hochreiter 1997. | |
| 714 | 791 | /// </summary> | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -17,6 +17,21 @@ public DropoutRNNCellMixin(LayerArgs args): base(args) | |||
| 17 | 17 | ||
| 18 | 18 | } | |
| 19 | 19 | ||
| 20 | + protected void _create_non_trackable_mask_cache() | ||
| 21 | + { | ||
| 22 | + | ||
| 23 | + } | ||
| 24 | + | ||
| 25 | + public void reset_dropout_mask() | ||
| 26 | + { | ||
| 27 | + | ||
| 28 | + } | ||
| 29 | + | ||
| 30 | + public void reset_recurrent_dropout_mask() | ||
| 31 | + { | ||
| 32 | + | ||
| 33 | + } | ||
| 34 | + | ||
| 20 | 35 | public Tensors? get_dropout_maskcell_for_cell(Tensors input, bool training, int count = 1) | |
| 21 | 36 | { | |
| 22 | 37 | if (dropout == 0f) | |
| Back | FazBrowse Home | New Git URL |
0 commit comments