| FazBrowse GitHub Viewer | Trending | | Home |
| Tools: [Download Repo ZIP] [Original HTTPS Page] |
1 parent 4c6063d commit 3805771
3 files changed
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -14,8 +14,10 @@ You may obtain a copy of the License at | |||
| 14 | 14 | limitations under the License. | |
| 15 | 15 | ******************************************************************************/ | |
| 16 | 16 | ||
| 17 | + using System.Xml.Linq; | ||
| 17 | 18 | using Tensorflow.Operations; | |
| 18 | 19 | using Tensorflow.Operations.Activation; | |
| 20 | + //using static System.Formats.Asn1.AsnWriter; | ||
| 19 | 21 | using static Tensorflow.Binding; | |
| 20 | 22 | ||
| 21 | 23 | namespace Tensorflow | |
@@ -125,6 +127,22 @@ public Tensor[] fused_batch_norm(Tensor x, | |||
| 125 | 127 | is_training: is_training, | |
| 126 | 128 | name: name, | |
| 127 | 129 | exponential_avg_factor: exponential_avg_factor); | |
| 130 | + public Tensor batch_normalization(Tensor x, | ||
| 131 | + Tensor mean, | ||
| 132 | + Tensor variance, | ||
| 133 | + Tensor offset, | ||
| 134 | + Tensor scale, | ||
| 135 | + float variance_epsilon, | ||
| 136 | + string name = null) | ||
| 137 | + { | ||
| 138 | + var inv = math_ops.rsqrt(variance + variance_epsilon); | ||
| 139 | + tf_with(ops.name_scope(name, "batchnorm", (x, mean, variance, scale, offset)), scope => | ||
| 140 | + { | ||
| 141 | + if (scale != null) inv *= scale; | ||
| 142 | + }); | ||
| 143 | + if (offset != null) return x * math_ops.cast(inv, x.dtype) + math_ops.cast(offset - mean * inv, dtype: x.dtype); | ||
| 144 | + else return x * math_ops.cast(inv, x.dtype) + math_ops.cast(-mean * inv, dtype: x.dtype); | ||
| 145 | + } | ||
| 128 | 146 | ||
| 129 | 147 | public Tensor max_pool(Tensor value, int[] ksize, int[] strides, string padding, string data_format = "NHWC", string name = null) | |
| 130 | 148 | => nn_ops.max_pool(value, ksize, strides, padding, data_format: data_format, name: name); | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -153,9 +153,22 @@ protected override Tensors Call(Tensors inputs, Tensors state = null, bool? trai | |||
| 153 | 153 | } | |
| 154 | 154 | else | |
| 155 | 155 | { | |
| 156 | + var input_dtype = inputs.dtype; | ||
| 157 | + if ((input_dtype == tf.float16) && DType == tf.float32) inputs = tf.cast(inputs, tf.float32); | ||
| 158 | + (Tensor mean, Tensor variance) = tf.nn.moments(inputs, axis, keep_dims: true); | ||
| 156 | 159 | ||
| 157 | - } | ||
| 160 | + (Tensor scale, Tensor offset) = (_broadcast(gamma), _broadcast(beta)); | ||
| 161 | + | ||
| 162 | + outputs = tf.nn.batch_normalization( | ||
| 163 | + inputs, | ||
| 164 | + mean, | ||
| 165 | + variance, | ||
| 166 | + offset: offset, | ||
| 167 | + scale: scale, | ||
| 168 | + variance_epsilon: epsilon); | ||
| 158 | 169 | ||
| 170 | + outputs = tf.cast(outputs, input_dtype); | ||
| 171 | + } | ||
| 159 | 172 | // If some components of the shape got lost due to adjustments, fix that. | |
| 160 | 173 | outputs.shape = input_shape; | |
| 161 | 174 | ||
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -1,5 +1,7 @@ | |||
| 1 | 1 | using Microsoft.VisualStudio.TestTools.UnitTesting; | |
| 2 | + using System; | ||
| 2 | 3 | using System.Collections.Generic; | |
| 4 | + using System.Linq; | ||
| 3 | 5 | using Tensorflow.NumPy; | |
| 4 | 6 | using static Tensorflow.Binding; | |
| 5 | 7 | using static Tensorflow.KerasApi; | |
@@ -161,6 +163,26 @@ public void LayerNormalization() | |||
| 161 | 163 | Tensor output = layer.Apply(inputs); | |
| 162 | 164 | Assert.AreEqual((5, 2), output.shape); | |
| 163 | 165 | Assert.IsTrue(output[0].numpy().Equals(new[] { -0.99998f, 0.99998f })); | |
| 166 | + | ||
| 167 | + // test_layernorm_weights | ||
| 168 | + Assert.AreEqual(len(layer.TrainableWeights), 2); | ||
| 169 | + Assert.AreEqual(len(layer.Weights), 2); | ||
| 170 | + | ||
| 171 | + var beta = layer.Weights.Where(x => x.Name.StartsWith("beta")).Single(); | ||
| 172 | + var gamma = layer.Weights.Where(x => x.Name.StartsWith("gamma")).Single(); | ||
| 173 | + | ||
| 174 | + // correctness_test | ||
| 175 | + layer = keras.layers.LayerNormalization(axis: -1, epsilon: (float) 1e-12); | ||
| 176 | + var x = np.random.normal(loc: 5.0f, scale: 10.0f, size: (1000, 2, 2, 2)).astype(tf.float32); | ||
| 177 | + | ||
| 178 | + output = layer.Apply(x); | ||
| 179 | + | ||
| 180 | + var y = (output - beta.numpy()) / gamma.numpy(); | ||
| 181 | + | ||
| 182 | + var y_mean = np.mean(y.numpy()); | ||
| 183 | + var y_std = np.sqrt(np.sum(np.power(y.numpy() - np.mean(y.numpy()), 2)) / 8000); | ||
| 184 | + Assert.IsTrue(tf.greater(np.array(0.1f), tf.abs(y_std - 1.0)).ToArray<bool>()[0]); | ||
| 185 | + Assert.IsTrue(tf.greater(np.array(0.1f), tf.abs(y_mean)).ToArray<bool>()[0]); | ||
| 164 | 186 | } | |
| 165 | 187 | ||
| 166 | 188 | /// <summary> | |
| Back | FazBrowse Home | New Git URL |
0 commit comments