| FazBrowse GitHub Viewer | Trending | | Home |
| Tools: [Download Repo ZIP] [Original HTTPS Page] |
1 parent 079b9a3 commit 4e42d7f
4 files changed
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -428,9 +428,9 @@ public static Tensor _transpose_batch_time(Tensor x) | |||
| 428 | 428 | return x; | |
| 429 | 429 | ||
| 430 | 430 | var x_rank = array_ops.rank(x); | |
| 431 | - var con1 = new object[] | ||
| 431 | + var con1 = new Tensor[] | ||
| 432 | 432 | { | |
| 433 | - new []{1, 0 }, | ||
| 433 | + new Tensor(new int[]{0, 2}), | ||
| 434 | 434 | math_ops.range(2, x_rank) | |
| 435 | 435 | }; | |
| 436 | 436 | var x_t = array_ops.transpose(x, array_ops.concat(con1, 0)); | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -166,6 +166,11 @@ public static Tensor boolean_mask<T1, T2>(T1 tensor, T2 mask, string name = "boo | |||
| 166 | 166 | throw new ValueError("mask cannot be scalar."); | |
| 167 | 167 | ||
| 168 | 168 | var leading_size = gen_math_ops.prod(shape(tensor_tensor)[$"{axis}:{axis + ndims_mask}"], ops.convert_to_tensor(new[] { 0 })); | |
| 169 | + if (leading_size.rank == 0) | ||
| 170 | + { | ||
| 171 | + leading_size = expand_dims(leading_size, 0); | ||
| 172 | + } | ||
| 173 | + | ||
| 169 | 174 | var shape1 = concat(new[] | |
| 170 | 175 | { | |
| 171 | 176 | shape(tensor_tensor)[$":{axis}"], | |
@@ -185,7 +190,7 @@ public static Tensor boolean_mask<T1, T2>(T1 tensor, T2 mask, string name = "boo | |||
| 185 | 190 | ||
| 186 | 191 | private static Tensor _apply_mask_1d(Tensor reshaped_tensor, Tensor mask, int axis = 0) | |
| 187 | 192 | { | |
| 188 | - var indices = squeeze(where(mask), axis: new[] { 1 }); | ||
| 193 | + var indices = squeeze(where_v2(mask), axis: new[] { 1 }); | ||
| 189 | 194 | return gather(reshaped_tensor, indices, axis: ops.convert_to_tensor(axis)); | |
| 190 | 195 | } | |
| 191 | 196 | ||
@@ -940,12 +945,12 @@ public static Tensor broadcast_static_shape(Tensor shape_x, Tensor shape_y) | |||
| 940 | 945 | /// <returns></returns> | |
| 941 | 946 | public static Tensor concat(Tensor[] values, Tensor axis, string name = "concat") | |
| 942 | 947 | { | |
| 943 | - return tf.Context.ExecuteOp("ConcatV2", name, new ExecuteOpArgs(values, axis)); | ||
| 948 | + return gen_array_ops.concat_v2(values, axis, name: name); | ||
| 944 | 949 | } | |
| 945 | 950 | ||
| 946 | - public static Tensor concat(object[] values, int axis, string name = "concat") | ||
| 951 | + public static Tensor concat(Tensor[] values, Axis axis, string name = "concat") | ||
| 947 | 952 | { | |
| 948 | - return tf.Context.ExecuteOp("ConcatV2", name, new ExecuteOpArgs(values, axis)); | ||
| 953 | + return gen_array_ops.concat_v2(values, axis, name: name); | ||
| 949 | 954 | } | |
| 950 | 955 | ||
| 951 | 956 | /// <summary> | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -287,7 +287,7 @@ private static Tensor _flatten_outer_dims(Tensor logits) | |||
| 287 | 287 | new[] { math_ops.subtract(rank, 1) }, | |
| 288 | 288 | new[] { constant_op.constant(1) }); | |
| 289 | 289 | ||
| 290 | - var ops = array_ops.concat(new[] { new[] { -1 }, (object)last_dim_size }, 0); | ||
| 290 | + var ops = array_ops.concat(new Tensor[] { new Tensor(new int[] {1}), last_dim_size }, 0); | ||
| 291 | 291 | var output = array_ops.reshape(logits, ops); | |
| 292 | 292 | ||
| 293 | 293 | // Set output shape if known. | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -3,6 +3,7 @@ | |||
| 3 | 3 | using System; | |
| 4 | 4 | using System.Linq; | |
| 5 | 5 | using static Tensorflow.Binding; | |
| 6 | + using Tensorflow; | ||
| 6 | 7 | ||
| 7 | 8 | namespace TensorFlowNET.UnitTest.Basics | |
| 8 | 9 | { | |
@@ -60,14 +61,14 @@ public void batch_to_space_nd() | |||
| 60 | 61 | Assert.IsTrue(Enumerable.SequenceEqual(new int[] { 15, 21, 16, 22, 17, 23 }, result[0, 3].ToArray<int>())); | |
| 61 | 62 | } | |
| 62 | 63 | ||
| 63 | - [TestMethod, Ignore] | ||
| 64 | + [TestMethod] | ||
| 64 | 65 | public void boolean_mask() | |
| 65 | 66 | { | |
| 67 | + if (!tf.executing_eagerly()) | ||
| 68 | + tf.enable_eager_execution(); | ||
| 66 | 69 | var tensor = new[] { 0, 1, 2, 3 }; | |
| 67 | 70 | var mask = np.array(new[] { true, false, true, false }); | |
| 68 | 71 | var masked = tf.boolean_mask(tensor, mask); | |
| 69 | - var sess = tf.Session(); | ||
| 70 | - var result = sess.run(masked); | ||
| 71 | 72 | Assert.IsTrue(Enumerable.SequenceEqual(new int[] { 0, 2 }, masked.ToArray<int>())); | |
| 72 | 73 | } | |
| 73 | 74 | } | |
| Back | FazBrowse Home | New Git URL |
0 commit comments