| FazBrowse GitHub Viewer | Trending | | Home |
| Tools: [Download Repo ZIP] [Original HTTPS Page] |
1 parent e72fd53 commit ee0b935
10 files changed
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -346,4 +346,5 @@ Get started with the implementation: | |||
| 346 | 346 | } | |
| 347 | 347 | ``` | |
| 348 | 348 | ||
| 349 | -  | ||
| 349 | +  | ||
| 350 | + | ||
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -16,9 +16,11 @@ limitations under the License. | |||
| 16 | 16 | ||
| 17 | 17 | using System; | |
| 18 | 18 | using System.Collections.Generic; | |
| 19 | + using System.Linq; | ||
| 19 | 20 | using System.Text; | |
| 20 | 21 | using Tensorflow.Operations; | |
| 21 | 22 | using Tensorflow.Operations.Activation; | |
| 23 | + using Tensorflow.Util; | ||
| 22 | 24 | using static Tensorflow.Python; | |
| 23 | 25 | ||
| 24 | 26 | namespace Tensorflow | |
@@ -68,6 +70,33 @@ public static Tensor dropout(Tensor x, Tensor keep_prob = null, Tensor noise_sha | |||
| 68 | 70 | return nn_ops.dropout_v2(x, rate: rate_tensor, noise_shape: noise_shape, seed: seed, name: name); | |
| 69 | 71 | } | |
| 70 | 72 | ||
| 73 | + /// <summary> | ||
| 74 | + /// Creates a recurrent neural network specified by RNNCell `cell`. | ||
| 75 | + /// </summary> | ||
| 76 | + /// <param name="cell">An instance of RNNCell.</param> | ||
| 77 | + /// <param name="inputs">The RNN inputs.</param> | ||
| 78 | + /// <param name="dtype"></param> | ||
| 79 | + /// <param name="swap_memory"></param> | ||
| 80 | + /// <param name="time_major"></param> | ||
| 81 | + /// <returns>A pair (outputs, state)</returns> | ||
| 82 | + public static (Tensor, Tensor) dynamic_rnn(RNNCell cell, Tensor inputs, TF_DataType dtype = TF_DataType.DtInvalid, | ||
| 83 | + bool swap_memory = false, bool time_major = false) | ||
| 84 | + { | ||
| 85 | + with(variable_scope("rnn"), scope => | ||
| 86 | + { | ||
| 87 | + VariableScope varscope = scope; | ||
| 88 | + var flat_input = nest.flatten(inputs); | ||
| 89 | + | ||
| 90 | + if (!time_major) | ||
| 91 | + { | ||
| 92 | + flat_input = flat_input.Select(x => ops.convert_to_tensor(x)).ToList(); | ||
| 93 | + //flat_input = flat_input.Select(x => _transpose_batch_time(x)).ToList(); | ||
| 94 | + } | ||
| 95 | + }); | ||
| 96 | + | ||
| 97 | + throw new NotImplementedException(""); | ||
| 98 | + } | ||
| 99 | + | ||
| 71 | 100 | public static (Tensor, Tensor) moments(Tensor x, | |
| 72 | 101 | int[] axes, | |
| 73 | 102 | string name = null, | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -1,10 +1,48 @@ | |||
| 1 | - using System; | ||
| 1 | + /***************************************************************************** | ||
| 2 | + Copyright 2018 The TensorFlow.NET Authors. All Rights Reserved. | ||
| 3 | + | ||
| 4 | + Licensed under the Apache License, Version 2.0 (the "License"); | ||
| 5 | + you may not use this file except in compliance with the License. | ||
| 6 | + You may obtain a copy of the License at | ||
| 7 | + | ||
| 8 | + http://www.apache.org/licenses/LICENSE-2.0 | ||
| 9 | + | ||
| 10 | + Unless required by applicable law or agreed to in writing, software | ||
| 11 | + distributed under the License is distributed on an "AS IS" BASIS, | ||
| 12 | + WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. | ||
| 13 | + See the License for the specific language governing permissions and | ||
| 14 | + limitations under the License. | ||
| 15 | + ******************************************************************************/ | ||
| 16 | + | ||
| 17 | + using System; | ||
| 2 | 18 | using System.Collections.Generic; | |
| 3 | 19 | using System.Text; | |
| 20 | + using Tensorflow.Keras.Engine; | ||
| 21 | + using Tensorflow.Operations.Activation; | ||
| 4 | 22 | ||
| 5 | 23 | namespace Tensorflow | |
| 6 | 24 | { | |
| 7 | - public class BasicRNNCell | ||
| 25 | + public class BasicRNNCell : LayerRNNCell | ||
| 8 | 26 | { | |
| 27 | + int _num_units; | ||
| 28 | + Func<Tensor, string, Tensor> _activation; | ||
| 29 | + | ||
| 30 | + public BasicRNNCell(int num_units, | ||
| 31 | + Func<Tensor, string, Tensor> activation = null, | ||
| 32 | + bool? reuse = null, | ||
| 33 | + string name = null, | ||
| 34 | + TF_DataType dtype = TF_DataType.DtInvalid) : base(_reuse: reuse, | ||
| 35 | + name: name, | ||
| 36 | + dtype: dtype) | ||
| 37 | + { | ||
| 38 | + // Inputs must be 2-dimensional. | ||
| 39 | + input_spec = new InputSpec(ndim: 2); | ||
| 40 | + | ||
| 41 | + _num_units = num_units; | ||
| 42 | + if (activation == null) | ||
| 43 | + _activation = math_ops.tanh; | ||
| 44 | + else | ||
| 45 | + _activation = activation; | ||
| 46 | + } | ||
| 9 | 47 | } | |
| 10 | 48 | } | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -0,0 +1,33 @@ | |||
| 1 | + /***************************************************************************** | ||
| 2 | + Copyright 2018 The TensorFlow.NET Authors. All Rights Reserved. | ||
| 3 | + | ||
| 4 | + Licensed under the Apache License, Version 2.0 (the "License"); | ||
| 5 | + you may not use this file except in compliance with the License. | ||
| 6 | + You may obtain a copy of the License at | ||
| 7 | + | ||
| 8 | + http://www.apache.org/licenses/LICENSE-2.0 | ||
| 9 | + | ||
| 10 | + Unless required by applicable law or agreed to in writing, software | ||
| 11 | + distributed under the License is distributed on an "AS IS" BASIS, | ||
| 12 | + WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. | ||
| 13 | + See the License for the specific language governing permissions and | ||
| 14 | + limitations under the License. | ||
| 15 | + ******************************************************************************/ | ||
| 16 | + | ||
| 17 | + using System; | ||
| 18 | + using System.Collections.Generic; | ||
| 19 | + using System.Text; | ||
| 20 | + | ||
| 21 | + namespace Tensorflow | ||
| 22 | + { | ||
| 23 | + public class LayerRNNCell : RNNCell | ||
| 24 | + { | ||
| 25 | + public LayerRNNCell(bool? _reuse = null, | ||
| 26 | + string name = null, | ||
| 27 | + TF_DataType dtype = TF_DataType.DtInvalid) : base(_reuse: _reuse, | ||
| 28 | + name: name, | ||
| 29 | + dtype: dtype) | ||
| 30 | + { | ||
| 31 | + } | ||
| 32 | + } | ||
| 33 | + } | ||
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -0,0 +1,63 @@ | |||
| 1 | + /***************************************************************************** | ||
| 2 | + Copyright 2018 The TensorFlow.NET Authors. All Rights Reserved. | ||
| 3 | + | ||
| 4 | + Licensed under the Apache License, Version 2.0 (the "License"); | ||
| 5 | + you may not use this file except in compliance with the License. | ||
| 6 | + You may obtain a copy of the License at | ||
| 7 | + | ||
| 8 | + http://www.apache.org/licenses/LICENSE-2.0 | ||
| 9 | + | ||
| 10 | + Unless required by applicable law or agreed to in writing, software | ||
| 11 | + distributed under the License is distributed on an "AS IS" BASIS, | ||
| 12 | + WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. | ||
| 13 | + See the License for the specific language governing permissions and | ||
| 14 | + limitations under the License. | ||
| 15 | + ******************************************************************************/ | ||
| 16 | + | ||
| 17 | + using System; | ||
| 18 | + using System.Collections.Generic; | ||
| 19 | + using System.Text; | ||
| 20 | + | ||
| 21 | + namespace Tensorflow | ||
| 22 | + { | ||
| 23 | + /// <summary> | ||
| 24 | + /// Abstract object representing an RNN cell. | ||
| 25 | + /// | ||
| 26 | + /// Every `RNNCell` must have the properties below and implement `call` with | ||
| 27 | + /// the signature `(output, next_state) = call(input, state)`. The optional | ||
| 28 | + /// third input argument, `scope`, is allowed for backwards compatibility | ||
| 29 | + /// purposes; but should be left off for new subclasses. | ||
| 30 | + /// | ||
| 31 | + /// This definition of cell differs from the definition used in the literature. | ||
| 32 | + /// In the literature, 'cell' refers to an object with a single scalar output. | ||
| 33 | + /// This definition refers to a horizontal array of such units. | ||
| 34 | + /// | ||
| 35 | + /// An RNN cell, in the most abstract setting, is anything that has | ||
| 36 | + /// a state and performs some operation that takes a matrix of inputs. | ||
| 37 | + /// This operation results in an output matrix with `self.output_size` columns. | ||
| 38 | + /// If `self.state_size` is an integer, this operation also results in a new | ||
| 39 | + /// state matrix with `self.state_size` columns. If `self.state_size` is a | ||
| 40 | + /// (possibly nested tuple of) TensorShape object(s), then it should return a | ||
| 41 | + /// matching structure of Tensors having shape `[batch_size].concatenate(s)` | ||
| 42 | + /// for each `s` in `self.batch_size`. | ||
| 43 | + /// </summary> | ||
| 44 | + public abstract class RNNCell : Layers.Layer | ||
| 45 | + { | ||
| 46 | + /// <summary> | ||
| 47 | + /// Attribute that indicates whether the cell is a TF RNN cell, due the slight | ||
| 48 | + /// difference between TF and Keras RNN cell. | ||
| 49 | + /// </summary> | ||
| 50 | + protected bool _is_tf_rnn_cell = false; | ||
| 51 | + | ||
| 52 | + public RNNCell(bool trainable = true, | ||
| 53 | + string name = null, | ||
| 54 | + TF_DataType dtype = TF_DataType.DtInvalid, | ||
| 55 | + bool? _reuse = null) : base(trainable: trainable, | ||
| 56 | + name: name, | ||
| 57 | + dtype: dtype, | ||
| 58 | + _reuse: _reuse) | ||
| 59 | + { | ||
| 60 | + _is_tf_rnn_cell = true; | ||
| 61 | + } | ||
| 62 | + } | ||
| 63 | + } | ||
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -551,6 +551,9 @@ public static Tensor conj(Tensor x, string name = null) | |||
| 551 | 551 | }); | |
| 552 | 552 | } | |
| 553 | 553 | ||
| 554 | + public static Tensor tanh(Tensor x, string name = null) | ||
| 555 | + => gen_math_ops.tanh(x, name); | ||
| 556 | + | ||
| 554 | 557 | public static Tensor truediv(Tensor x, Tensor y, string name = null) | |
| 555 | 558 | => _truediv_python3(x, y, name); | |
| 556 | 559 | ||
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -7,8 +7,6 @@ namespace Tensorflow.Operations | |||
| 7 | 7 | public class rnn_cell_impl | |
| 8 | 8 | { | |
| 9 | 9 | public BasicRNNCell BasicRNNCell(int num_units) | |
| 10 | - { | ||
| 11 | - throw new NotImplementedException(); | ||
| 12 | - } | ||
| 10 | + => new BasicRNNCell(num_units); | ||
| 13 | 11 | } | |
| 14 | 12 | } | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -214,14 +214,14 @@ public static bool is_sequence(object arg) | |||
| 214 | 214 | //# See the swig file (util.i) for documentation. | |
| 215 | 215 | //flatten = _pywrap_tensorflow.Flatten | |
| 216 | 216 | ||
| 217 | - public static List<object> flatten(object structure) | ||
| 217 | + public static List<T> flatten<T>(T structure) | ||
| 218 | 218 | { | |
| 219 | - var list = new List<object>(); | ||
| 219 | + var list = new List<T>(); | ||
| 220 | 220 | _flatten_recursive(structure, list); | |
| 221 | 221 | return list; | |
| 222 | 222 | } | |
| 223 | 223 | ||
| 224 | - private static void _flatten_recursive(object obj, List<object> list) | ||
| 224 | + private static void _flatten_recursive<T>(T obj, List<T> list) | ||
| 225 | 225 | { | |
| 226 | 226 | if (obj is string) | |
| 227 | 227 | { | |
@@ -232,7 +232,7 @@ private static void _flatten_recursive(object obj, List<object> list) | |||
| 232 | 232 | { | |
| 233 | 233 | var dict = obj as IDictionary; | |
| 234 | 234 | foreach (var key in _sorted(dict)) | |
| 235 | - _flatten_recursive(dict[key], list); | ||
| 235 | + _flatten_recursive((T)dict[key], list); | ||
| 236 | 236 | return; | |
| 237 | 237 | } | |
| 238 | 238 | if (obj is NDArray) | |
@@ -244,7 +244,7 @@ private static void _flatten_recursive(object obj, List<object> list) | |||
| 244 | 244 | { | |
| 245 | 245 | var structure = obj as IEnumerable; | |
| 246 | 246 | foreach (var child in structure) | |
| 247 | - _flatten_recursive(child, list); | ||
| 247 | + _flatten_recursive((T)child, list); | ||
| 248 | 248 | return; | |
| 249 | 249 | } | |
| 250 | 250 | list.Add(obj); | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -25,14 +25,12 @@ limitations under the License. | |||
| 25 | 25 | namespace TensorFlowNET.Examples.ImageProcess | |
| 26 | 26 | { | |
| 27 | 27 | /// <summary> | |
| 28 | - /// Convolutional Neural Network classifier for Hand Written Digits | ||
| 29 | - /// CNN architecture with two convolutional layers, followed by two fully-connected layers at the end. | ||
| 30 | - /// Use Stochastic Gradient Descent (SGD) optimizer. | ||
| 31 | - /// http://www.easy-tensorflow.com/tf-tutorials/convolutional-neural-nets-cnns/cnn1 | ||
| 28 | + /// Recurrent Neural Network for handwritten digits MNIST. | ||
| 29 | + /// https://medium.com/machine-learning-algorithms/mnist-using-recurrent-neural-network-2d070a5915a2 | ||
| 32 | 30 | /// </summary> | |
| 33 | 31 | public class DigitRecognitionRNN : IExample | |
| 34 | 32 | { | |
| 35 | - public bool Enabled { get; set; } = false; | ||
| 33 | + public bool Enabled { get; set; } = true; | ||
| 36 | 34 | public bool IsImportingGraph { get; set; } = false; | |
| 37 | 35 | ||
| 38 | 36 | public string Name => "MNIST RNN"; | |
@@ -84,6 +82,7 @@ public Graph BuildGraph() | |||
| 84 | 82 | var X = tf.placeholder(tf.float32, new[] { -1, n_steps, n_inputs }); | |
| 85 | 83 | var y = tf.placeholder(tf.int32, new[] { -1 }); | |
| 86 | 84 | var cell = tf.nn.rnn_cell.BasicRNNCell(num_units: n_neurons); | |
| 85 | + var (output, state) = tf.nn.dynamic_rnn(cell, X, dtype: tf.float32); | ||
| 87 | 86 | ||
| 88 | 87 | return graph; | |
| 89 | 88 | } | |
@@ -154,6 +153,7 @@ public void PrepareData() | |||
| 154 | 153 | print("Size of:"); | |
| 155 | 154 | print($"- Training-set:\t\t{len(mnist.train.data)}"); | |
| 156 | 155 | print($"- Validation-set:\t{len(mnist.validation.data)}"); | |
| 156 | + print($"- Test-set:\t\t{len(mnist.test.data)}"); | ||
| 157 | 157 | } | |
| 158 | 158 | ||
| 159 | 159 | public Graph ImportGraph() => throw new NotImplementedException(); | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -78,8 +78,8 @@ public void testFlattenAndPack() | |||
| 78 | 78 | self.assertEqual((((restructured_from_flat[1] as object[])[0] as object[])[0] as Hashtable)["y"], 0); | |
| 79 | 79 | ||
| 80 | 80 | self.assertEqual(new List<object> { 5 }, nest.flatten(5)); | |
| 81 | - flat = nest.flatten(np.array(new[] { 5 })); | ||
| 82 | - self.assertEqual(new object[] { np.array(new int[] { 5 }) }, flat); | ||
| 81 | + var flat1 = nest.flatten(np.array(new[] { 5 })); | ||
| 82 | + self.assertEqual(new object[] { np.array(new int[] { 5 }) }, flat1); | ||
| 83 | 83 | ||
| 84 | 84 | self.assertEqual("a", nest.pack_sequence_as(5, new List<object> { "a" })); | |
| 85 | 85 | self.assertEqual(np.array(new[] { 5 }), | |
| Back | FazBrowse Home | New Git URL |
0 commit comments