| FazBrowse GitHub Viewer | Trending | | Home |
| Tools: [Download Repo ZIP] [Original HTTPS Page] |
1 parent cfffc68 commit 6a8665f
35 files changed
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -149,13 +149,18 @@ public static int len(object a) | |||
| 149 | 149 | return ndArray.ndim == 0 ? 1 : ndArray.shape[0]; | |
| 150 | 150 | case IEnumerable enumerable: | |
| 151 | 151 | return enumerable.OfType<object>().Count(); | |
| 152 | + case TensorShape arr: | ||
| 153 | + return arr.ndim; | ||
| 152 | 154 | } | |
| 153 | 155 | throw new NotImplementedException("len() not implemented for type: " + a.GetType()); | |
| 154 | 156 | } | |
| 155 | 157 | ||
| 156 | 158 | public static float min(float a, float b) | |
| 157 | 159 | => Math.Min(a, b); | |
| 158 | 160 | ||
| 161 | + public static int max(int a, int b) | ||
| 162 | + => Math.Max(a, b); | ||
| 163 | + | ||
| 159 | 164 | public static T[] list<T>(IEnumerable<T> list) | |
| 160 | 165 | => list.ToArray(); | |
| 161 | 166 | ||
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -15,6 +15,8 @@ limitations under the License. | |||
| 15 | 15 | ******************************************************************************/ | |
| 16 | 16 | ||
| 17 | 17 | using System; | |
| 18 | + using System.Linq; | ||
| 19 | + using static Tensorflow.Binding; | ||
| 18 | 20 | ||
| 19 | 21 | namespace Tensorflow.Framework | |
| 20 | 22 | { | |
@@ -52,7 +54,14 @@ public static Tensor smart_cond(bool pred, | |||
| 52 | 54 | { | |
| 53 | 55 | var pred_value = tensor_util.constant_value(pred); | |
| 54 | 56 | if (pred_value is null) | |
| 55 | - return pred.eval(new Session(pred.graph)); | ||
| 57 | + { | ||
| 58 | + var result = range(pred.op.NumOutputs).Select(x => IntPtr.Zero).ToArray(); | ||
| 59 | + var evaluated = c_api.TF_TryEvaluateConstant(pred.graph, pred._as_tf_output(), result, tf.Status.Handle); | ||
| 60 | + if (!evaluated || c_api.TF_GetCode(tf.Status.Handle) != TF_Code.TF_OK) | ||
| 61 | + return null; | ||
| 62 | + else | ||
| 63 | + throw new NotImplementedException(""); | ||
| 64 | + } | ||
| 56 | 65 | ||
| 57 | 66 | return pred_value; | |
| 58 | 67 | } | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -322,5 +322,18 @@ public static extern void TF_GraphSetOutputHandleShapesAndTypes(IntPtr graph, TF | |||
| 322 | 322 | [DllImport(TensorFlowLibName)] | |
| 323 | 323 | ||
| 324 | 324 | public static extern void TF_UpdateEdge(IntPtr graph, TF_Output new_src, TF_Input dst, SafeStatusHandle status); | |
| 325 | + | ||
| 326 | + /// <summary> | ||
| 327 | + /// Attempts to evaluate `output`. This will only be possible if `output` doesn't | ||
| 328 | + /// depend on any graph inputs (this function is safe to call if this isn't the | ||
| 329 | + /// case though). | ||
| 330 | + /// </summary> | ||
| 331 | + /// <param name="graph"></param> | ||
| 332 | + /// <param name="output"></param> | ||
| 333 | + /// <param name="result"></param> | ||
| 334 | + /// <param name="status"></param> | ||
| 335 | + /// <returns></returns> | ||
| 336 | + [DllImport(TensorFlowLibName)] | ||
| 337 | + public static extern bool TF_TryEvaluateConstant(IntPtr graph, TF_Output output, IntPtr[] result, SafeStatusHandle status); | ||
| 325 | 338 | } | |
| 326 | 339 | } | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -50,6 +50,6 @@ public InputSpec(TF_DataType dtype = TF_DataType.DtInvalid, | |||
| 50 | 50 | } | |
| 51 | 51 | ||
| 52 | 52 | public override string ToString() | |
| 53 | - => $"min_ndim={min_ndim}, , axes={axes.Count}"; | ||
| 53 | + => $"ndim={ndim}, min_ndim={min_ndim}, axes={axes.Count}"; | ||
| 54 | 54 | } | |
| 55 | 55 | } | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -21,6 +21,31 @@ namespace Tensorflow | |||
| 21 | 21 | { | |
| 22 | 22 | public class nn_impl | |
| 23 | 23 | { | |
| 24 | + public static Tensor conv2d_transpose(Tensor value = null, | ||
| 25 | + IVariableV1 filter = null, | ||
| 26 | + Tensor output_shape = null, | ||
| 27 | + TensorShape strides = null, | ||
| 28 | + string padding = "SAME", | ||
| 29 | + string data_format = "NHWC", | ||
| 30 | + string name = null, | ||
| 31 | + TensorShape dilations = null) | ||
| 32 | + { | ||
| 33 | + if (dilations == null) | ||
| 34 | + dilations = (1, 1, 1, 1); | ||
| 35 | + return tf_with(ops.name_scope(name, "conv2d_transpose", new { value, filter, output_shape }), scope => | ||
| 36 | + { | ||
| 37 | + return gen_nn_ops.conv2d_backprop_input( | ||
| 38 | + input_sizes: output_shape, | ||
| 39 | + filter: filter.AsTensor(), | ||
| 40 | + out_backprop: value, | ||
| 41 | + strides: strides, | ||
| 42 | + padding: padding, | ||
| 43 | + data_format: data_format, | ||
| 44 | + dilations: dilations, | ||
| 45 | + name: name); | ||
| 46 | + }); | ||
| 47 | + } | ||
| 48 | + | ||
| 24 | 49 | /// <summary> | |
| 25 | 50 | /// Normalizes along dimension `axis` using an L2 norm. | |
| 26 | 51 | /// </summary> | |
@@ -83,6 +108,23 @@ public static (Tensor, Tensor) moments(Tensor x, | |||
| 83 | 108 | }); | |
| 84 | 109 | } | |
| 85 | 110 | ||
| 111 | + public static Tensor batch_normalization(Tensor x, | ||
| 112 | + Tensor mean, | ||
| 113 | + Tensor variance, | ||
| 114 | + Tensor offset, | ||
| 115 | + Tensor scale, | ||
| 116 | + float variance_epsilon = 0.001f, | ||
| 117 | + string name = null) | ||
| 118 | + { | ||
| 119 | + return tf_with(ops.name_scope(name, "batchnorm", new { x, mean, variance, scale, offset }), scope => | ||
| 120 | + { | ||
| 121 | + var inv = math_ops.rsqrt(variance + variance_epsilon); | ||
| 122 | + inv *= scale; | ||
| 123 | + return x * math_ops.cast(inv, x.dtype) + math_ops.cast( | ||
| 124 | + offset == null ? (-mean * inv) : (offset - mean * inv), x.dtype); | ||
| 125 | + }); | ||
| 126 | + } | ||
| 127 | + | ||
| 86 | 128 | /// <summary> | |
| 87 | 129 | /// Batch normalization. | |
| 88 | 130 | /// </summary> | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -15,6 +15,10 @@ public override bool Equals(Object obj) | |||
| 15 | 15 | else if (rank != shape1.rank) | |
| 16 | 16 | return false; | |
| 17 | 17 | return Enumerable.SequenceEqual(shape1.dims, dims); | |
| 18 | + case int[] shape2: | ||
| 19 | + if (rank != shape2.Length) | ||
| 20 | + return false; | ||
| 21 | + return Enumerable.SequenceEqual(dims, shape2); | ||
| 18 | 22 | default: | |
| 19 | 23 | return false; | |
| 20 | 24 | } | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -317,5 +317,30 @@ public Tensor concatenate(Tensors tensors, int axis = -1) | |||
| 317 | 317 | ||
| 318 | 318 | return array_ops.concat(tensors, axis); | |
| 319 | 319 | } | |
| 320 | + | ||
| 321 | + public Tensor conv2d_transpose(Tensor x, | ||
| 322 | + IVariableV1 kernel, | ||
| 323 | + Tensor output_shape, | ||
| 324 | + TensorShape strides = null, | ||
| 325 | + string padding = "valid", | ||
| 326 | + string data_format = null, | ||
| 327 | + TensorShape dilation_rate = null) | ||
| 328 | + { | ||
| 329 | + var force_transpose = false; | ||
| 330 | + if (data_format == "channels_first" && !dilation_rate.Equals(new[] { 1, 1 })) | ||
| 331 | + force_transpose = true; | ||
| 332 | + // x, tf_data_format = _preprocess_conv2d_input(x, data_format, force_transpose) | ||
| 333 | + var tf_data_format = "NHWC"; | ||
| 334 | + padding = padding.ToUpper(); | ||
| 335 | + strides = new TensorShape(1, strides[0], strides[1], 1); | ||
| 336 | + if (dilation_rate.Equals(new[] { 1, 1 })) | ||
| 337 | + x = nn_impl.conv2d_transpose(x, kernel, output_shape, strides, | ||
| 338 | + padding: padding, | ||
| 339 | + data_format: tf_data_format); | ||
| 340 | + else | ||
| 341 | + throw new NotImplementedException(""); | ||
| 342 | + | ||
| 343 | + return x; | ||
| 344 | + } | ||
| 320 | 345 | } | |
| 321 | 346 | } | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -301,9 +301,9 @@ void BuildMapHelper(Tensor tensor, | |||
| 301 | 301 | nodes_in_decreasing_depth.append(node); | |
| 302 | 302 | } | |
| 303 | 303 | ||
| 304 | - protected override Tensors Call(Tensors inputs, Tensor state = null, bool is_training = false) | ||
| 304 | + protected override Tensors Call(Tensors inputs, Tensor state = null, bool? training = null) | ||
| 305 | 305 | { | |
| 306 | - return run_internal_graph(inputs, is_training); | ||
| 306 | + return run_internal_graph(inputs, training.Value); | ||
| 307 | 307 | } | |
| 308 | 308 | ||
| 309 | 309 | Tensors run_internal_graph(Tensors inputs, bool training = false, Tensors mask = null) | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -10,9 +10,9 @@ public partial class Layer | |||
| 10 | 10 | /// </summary> | |
| 11 | 11 | /// <param name="input"></param> | |
| 12 | 12 | /// <param name="state"></param> | |
| 13 | - /// <param name="is_training"></param> | ||
| 13 | + /// <param name="training"></param> | ||
| 14 | 14 | /// <returns></returns> | |
| 15 | - public Tensors Apply(Tensors inputs, Tensor state = null, bool is_training = false) | ||
| 15 | + public Tensors Apply(Tensors inputs, Tensor state = null, bool training = false) | ||
| 16 | 16 | { | |
| 17 | 17 | callContext = callContext ?? new ThreadLocal<CallContext>() | |
| 18 | 18 | { | |
@@ -38,7 +38,7 @@ public Tensors Apply(Tensors inputs, Tensor state = null, bool is_training = fal | |||
| 38 | 38 | if (!built) | |
| 39 | 39 | MaybeBuild(inputs); | |
| 40 | 40 | ||
| 41 | - outputs = Call(inputs, state: state, is_training: is_training); | ||
| 41 | + outputs = Call(inputs, state: state, training: training); | ||
| 42 | 42 | ||
| 43 | 43 | // memory leak | |
| 44 | 44 | // _set_connectivity_metadata_(inputs, outputs); | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -155,7 +155,7 @@ private Tensor compute_mask(Tensor inputs, Tensor mask = null) | |||
| 155 | 155 | /// <param name="state"></param> | |
| 156 | 156 | /// <param name="is_training"></param> | |
| 157 | 157 | /// <returns></returns> | |
| 158 | - protected virtual Tensors Call(Tensors inputs, Tensor state = null, bool is_training = false) | ||
| 158 | + protected virtual Tensors Call(Tensors inputs, Tensor state = null, bool? training = null) | ||
| 159 | 159 | { | |
| 160 | 160 | throw new NotImplementedException(""); | |
| 161 | 161 | } | |
| Back | FazBrowse Home | New Git URL |
0 commit comments