|
|
@@ -52,5 +52,19 @@ namespace Tensorflow.UnitTest.TF_API |
|
|
|
var concatValue = tf.concat(new[] { a, b, c }, axis: 0); |
|
|
|
Assert.IsTrue(Enumerable.SequenceEqual(new[] { 6, 2 }, concatValue.shape)); |
|
|
|
} |
|
|
|
[TestMethod] |
|
|
|
public void SplitTest() |
|
|
|
{ |
|
|
|
var a = tf.constant(new[,] { { 1, 2 }, { 3, 4 } }); |
|
|
|
var b = tf.constant(new[,] { { 5, 6 }, { 7, 8 } }); |
|
|
|
var c = tf.constant(new[,] { { 9, 10 }, { 11, 12 } }); |
|
|
|
|
|
|
|
var concatValue = tf.concat(new[] { a, b, c }, axis: 0); |
|
|
|
|
|
|
|
var splitValue = tf.split(concatValue, 3, axis: new Tensor(0)); |
|
|
|
Assert.IsTrue(Enumerable.SequenceEqual(new[] { 2, 2 }, splitValue[0].shape)); |
|
|
|
|
|
|
|
} |
|
|
|
|
|
|
|
} |
|
|
|
} |