diff --git a/tensorflow-examples/pom.xml b/tensorflow-examples/pom.xml index b6e41a7..bbc0adb 100644 --- a/tensorflow-examples/pom.xml +++ b/tensorflow-examples/pom.xml @@ -2,7 +2,7 @@ 4.0.0 org.tensorflow.model tensorflow-examples - 0.1.0-SNAPSHOT + 0.3.1-SNAPSHOT TensorFlow Examples A suite of executable examples using TensorFlow Java @@ -18,12 +18,12 @@ org.tensorflow tensorflow-core-platform - 0.2.0 + 0.3.1 org.tensorflow tensorflow-framework - 0.2.0 + 0.3.1 diff --git a/tensorflow-examples/src/main/java/org/tensorflow/model/examples/cnn/lenet/CnnMnist.java b/tensorflow-examples/src/main/java/org/tensorflow/model/examples/cnn/lenet/CnnMnist.java index 5b29722..fedd86d 100644 --- a/tensorflow-examples/src/main/java/org/tensorflow/model/examples/cnn/lenet/CnnMnist.java +++ b/tensorflow-examples/src/main/java/org/tensorflow/model/examples/cnn/lenet/CnnMnist.java @@ -88,22 +88,22 @@ public static Graph build(String optimizerName) { Ops tf = Ops.create(graph); // Inputs - Placeholder input = tf.withName(INPUT_NAME).placeholder(TUint8.DTYPE, + Placeholder input = tf.withName(INPUT_NAME).placeholder(TUint8.class, Placeholder.shape(Shape.of(-1, IMAGE_SIZE, IMAGE_SIZE))); Reshape input_reshaped = tf .reshape(input, tf.array(-1, IMAGE_SIZE, IMAGE_SIZE, NUM_CHANNELS)); - Placeholder labels = tf.withName(TARGET).placeholder(TUint8.DTYPE); + Placeholder labels = tf.withName(TARGET).placeholder(TUint8.class); // Scaling the features Constant centeringFactor = tf.constant(PIXEL_DEPTH / 2.0f); Constant scalingFactor = tf.constant((float) PIXEL_DEPTH); Operand scaledInput = tf.math - .div(tf.math.sub(tf.dtypes.cast(input_reshaped, TFloat32.DTYPE), centeringFactor), + .div(tf.math.sub(tf.dtypes.cast(input_reshaped, TFloat32.class), centeringFactor), scalingFactor); // First conv layer Variable conv1Weights = tf.variable(tf.math.mul(tf.random - .truncatedNormal(tf.array(5, 5, NUM_CHANNELS, 32), TFloat32.DTYPE, + .truncatedNormal(tf.array(5, 5, NUM_CHANNELS, 32), TFloat32.class, TruncatedNormal.seed(SEED)), tf.constant(0.1f))); Conv2d conv1 = tf.nn .conv2d(scaledInput, conv1Weights, Arrays.asList(1L, 1L, 1L, 1L), PADDING_TYPE); @@ -118,7 +118,7 @@ public static Graph build(String optimizerName) { // Second conv layer Variable conv2Weights = tf.variable(tf.math.mul(tf.random - .truncatedNormal(tf.array(5, 5, 32, 64), TFloat32.DTYPE, + .truncatedNormal(tf.array(5, 5, 32, 64), TFloat32.class, TruncatedNormal.seed(SEED)), tf.constant(0.1f))); Conv2d conv2 = tf.nn .conv2d(pool1, conv2Weights, Arrays.asList(1L, 1L, 1L, 1L), PADDING_TYPE); @@ -138,7 +138,7 @@ public static Graph build(String optimizerName) { // Fully connected layer Variable fc1Weights = tf.variable(tf.math.mul(tf.random - .truncatedNormal(tf.array(IMAGE_SIZE * IMAGE_SIZE * 4, 512), TFloat32.DTYPE, + .truncatedNormal(tf.array(IMAGE_SIZE * IMAGE_SIZE * 4, 512), TFloat32.class, TruncatedNormal.seed(SEED)), tf.constant(0.1f))); Variable fc1Biases = tf .variable(tf.fill(tf.array(new int[]{512}), tf.constant(0.1f))); @@ -147,7 +147,7 @@ public static Graph build(String optimizerName) { // Softmax layer Variable fc2Weights = tf.variable(tf.math.mul(tf.random - .truncatedNormal(tf.array(512, NUM_LABELS), TFloat32.DTYPE, + .truncatedNormal(tf.array(512, NUM_LABELS), TFloat32.class, TruncatedNormal.seed(SEED)), tf.constant(0.1f))); Variable fc2Biases = tf .variable(tf.fill(tf.array(new int[]{NUM_LABELS}), tf.constant(0.1f))); @@ -214,17 +214,17 @@ public static void train(Session session, int epochs, int minibatchSize, MnistDa // Train the model for (int i = 0; i < epochs; i++) { for (ImageBatch trainingBatch : dataset.trainingBatches(minibatchSize)) { - try (Tensor batchImages = TUint8.tensorOf(trainingBatch.images()); - Tensor batchLabels = TUint8.tensorOf(trainingBatch.labels()); - Tensor loss = session.runner() + try (TUint8 batchImages = TUint8.tensorOf(trainingBatch.images()); + TUint8 batchLabels = TUint8.tensorOf(trainingBatch.labels()); + TFloat32 loss = (TFloat32)session.runner() .feed(TARGET, batchLabels) .feed(INPUT_NAME, batchImages) .addTarget(TRAIN) .fetch(TRAINING_LOSS) - .run().get(0).expect(TFloat32.DTYPE)) { + .run().get(0)) { if (interval % 100 == 0) { logger.log(Level.INFO, - "Iteration = " + interval + ", training loss = " + loss.data().getFloat()); + "Iteration = " + interval + ", training loss = " + loss.getFloat()); } } interval++; @@ -237,17 +237,17 @@ public static void test(Session session, int minibatchSize, MnistDataset dataset int[][] confusionMatrix = new int[10][10]; for (ImageBatch trainingBatch : dataset.testBatches(minibatchSize)) { - try (Tensor transformedInput = TUint8.tensorOf(trainingBatch.images()); - Tensor outputTensor = session.runner() + try (TUint8 transformedInput = TUint8.tensorOf(trainingBatch.images()); + TFloat32 outputTensor = (TFloat32)session.runner() .feed(INPUT_NAME, transformedInput) - .fetch(OUTPUT_NAME).run().get(0).expect(TFloat32.DTYPE)) { + .fetch(OUTPUT_NAME).run().get(0)) { ByteNdArray labelBatch = trainingBatch.labels(); for (int k = 0; k < labelBatch.shape().size(0); k++) { byte trueLabel = labelBatch.getByte(k); int predLabel; - predLabel = argmax(outputTensor.data().slice(Indices.at(k), Indices.all())); + predLabel = argmax(outputTensor.slice(Indices.at(k), Indices.all())); if (predLabel == trueLabel) { correctCount++; } diff --git a/tensorflow-examples/src/main/java/org/tensorflow/model/examples/cnn/vgg/VGGModel.java b/tensorflow-examples/src/main/java/org/tensorflow/model/examples/cnn/vgg/VGGModel.java index 13f3ff6..9e9725d 100644 --- a/tensorflow-examples/src/main/java/org/tensorflow/model/examples/cnn/vgg/VGGModel.java +++ b/tensorflow-examples/src/main/java/org/tensorflow/model/examples/cnn/vgg/VGGModel.java @@ -83,17 +83,17 @@ public static Graph compile() { Ops tf = Ops.create(graph); // Inputs - Placeholder input = tf.withName(INPUT_NAME).placeholder(TUint8.DTYPE, + Placeholder input = tf.withName(INPUT_NAME).placeholder(TUint8.class, Placeholder.shape(Shape.of(-1, IMAGE_SIZE, IMAGE_SIZE))); Reshape input_reshaped = tf .reshape(input, tf.array(-1, IMAGE_SIZE, IMAGE_SIZE, NUM_CHANNELS)); - Placeholder labels = tf.withName(TARGET).placeholder(TUint8.DTYPE); + Placeholder labels = tf.withName(TARGET).placeholder(TUint8.class); // Scaling the features Constant centeringFactor = tf.constant(PIXEL_DEPTH / 2.0f); Constant scalingFactor = tf.constant((float) PIXEL_DEPTH); Operand scaledInput = tf.math - .div(tf.math.sub(tf.dtypes.cast(input_reshaped, TFloat32.DTYPE), centeringFactor), + .div(tf.math.sub(tf.dtypes.cast(input_reshaped, TFloat32.class), centeringFactor), scalingFactor); Relu relu1 = vggConv2DLayer("1", tf, scaledInput, new int[]{3, 3, NUM_CHANNELS, 32}, 32); @@ -137,7 +137,7 @@ public static Add buildFCLayersAndRegularization(Ops tf, Placeholder fc1Weights = tf.variable(tf.math.mul(tf.random - .truncatedNormal(tf.array(fcWeightShape), TFloat32.DTYPE, + .truncatedNormal(tf.array(fcWeightShape), TFloat32.class, TruncatedNormal.seed(SEED)), tf.constant(0.1f))); Variable fc1Biases = tf .variable(tf.fill(tf.array(new int[]{fcBiasShape}), tf.constant(0.1f))); @@ -146,7 +146,7 @@ public static Add buildFCLayersAndRegularization(Ops tf, Placeholder fc2Weights = tf.variable(tf.math.mul(tf.random - .truncatedNormal(tf.array(fcBiasShape, NUM_LABELS), TFloat32.DTYPE, + .truncatedNormal(tf.array(fcBiasShape, NUM_LABELS), TFloat32.class, TruncatedNormal.seed(SEED)), tf.constant(0.1f))); Variable fc2Biases = tf .variable(tf.fill(tf.array(new int[]{NUM_LABELS}), tf.constant(0.1f))); @@ -183,7 +183,7 @@ public static MaxPool vggMaxPool(Ops tf, Relu relu1) { public static Relu vggConv2DLayer(String layerName, Ops tf, Operand scaledInput, int[] convWeightsL1Shape, int convBiasL1Shape) { Variable conv1Weights = tf.withName("conv2d_" + layerName).variable(tf.math.mul(tf.random - .truncatedNormal(tf.array(convWeightsL1Shape), TFloat32.DTYPE, + .truncatedNormal(tf.array(convWeightsL1Shape), TFloat32.class, TruncatedNormal.seed(SEED)), tf.constant(0.1f))); Conv2d conv = tf.nn .conv2d(scaledInput, conv1Weights, Arrays.asList(1L, 1L, 1L, 1L), PADDING_TYPE); @@ -201,17 +201,17 @@ public void train(MnistDataset dataset, int epochs, int minibatchSize) { // Train the model for (int i = 0; i < epochs; i++) { for (ImageBatch trainingBatch : dataset.trainingBatches(minibatchSize)) { - try (Tensor batchImages = TUint8.tensorOf(trainingBatch.images()); - Tensor batchLabels = TUint8.tensorOf(trainingBatch.labels()); - Tensor loss = session.runner() + try (TUint8 batchImages = TUint8.tensorOf(trainingBatch.images()); + TUint8 batchLabels = TUint8.tensorOf(trainingBatch.labels()); + TFloat32 loss = (TFloat32)session.runner() .feed(TARGET, batchLabels) .feed(INPUT_NAME, batchImages) .addTarget(TRAIN) .fetch(TRAINING_LOSS) - .run().get(0).expect(TFloat32.DTYPE)) { + .run().get(0)) { logger.log(Level.INFO, - "Iteration = " + interval + ", training loss = " + loss.data().getFloat()); + "Iteration = " + interval + ", training loss = " + loss.getFloat()); } interval++; @@ -224,17 +224,17 @@ public void test(MnistDataset dataset, int minibatchSize) { int[][] confusionMatrix = new int[10][10]; for (ImageBatch trainingBatch : dataset.testBatches(minibatchSize)) { - try (Tensor transformedInput = TUint8.tensorOf(trainingBatch.images()); - Tensor outputTensor = session.runner() + try (TUint8 transformedInput = TUint8.tensorOf(trainingBatch.images()); + TFloat32 outputTensor = (TFloat32)session.runner() .feed(INPUT_NAME, transformedInput) - .fetch(OUTPUT_NAME).run().get(0).expect(TFloat32.DTYPE)) { + .fetch(OUTPUT_NAME).run().get(0)) { ByteNdArray labelBatch = trainingBatch.labels(); for (int k = 0; k < labelBatch.shape().size(0); k++) { byte trueLabel = labelBatch.getByte(k); int predLabel; - predLabel = argmax(outputTensor.data().slice(Indices.at(k), Indices.all())); + predLabel = argmax(outputTensor.slice(Indices.at(k), Indices.all())); if (predLabel == trueLabel) { correctCount++; } diff --git a/tensorflow-examples/src/main/java/org/tensorflow/model/examples/datasets/mnist/MnistDataset.java b/tensorflow-examples/src/main/java/org/tensorflow/model/examples/datasets/mnist/MnistDataset.java index a2ffe92..8c79ee4 100644 --- a/tensorflow-examples/src/main/java/org/tensorflow/model/examples/datasets/mnist/MnistDataset.java +++ b/tensorflow-examples/src/main/java/org/tensorflow/model/examples/datasets/mnist/MnistDataset.java @@ -27,8 +27,8 @@ import java.io.IOException; import java.util.zip.GZIPInputStream; -import static org.tensorflow.ndarray.index.Indices.from; -import static org.tensorflow.ndarray.index.Indices.to; +import static org.tensorflow.ndarray.index.Indices.sliceFrom; +import static org.tensorflow.ndarray.index.Indices.sliceTo; /** Common loader and data preprocessor for MNIST and FashionMNIST datasets. */ public class MnistDataset { @@ -44,10 +44,10 @@ public static MnistDataset create(int validationSize, String trainingImagesArchi if (validationSize > 0) { return new MnistDataset( - trainingImages.slice(from(validationSize)), - trainingLabels.slice(from(validationSize)), - trainingImages.slice(to(validationSize)), - trainingLabels.slice(to(validationSize)), + trainingImages.slice(sliceFrom(validationSize)), + trainingLabels.slice(sliceFrom(validationSize)), + trainingImages.slice(sliceTo(validationSize)), + trainingLabels.slice(sliceTo(validationSize)), testImages, testLabels ); @@ -137,6 +137,6 @@ private static ByteNdArray readArchive(String archiveName) throws IOException { } byte[] bytes = new byte[size]; archiveStream.readFully(bytes); - return NdArrays.wrap(DataBuffers.of(bytes, true, false), Shape.of(dimSizes)); + return NdArrays.wrap(Shape.of(dimSizes), DataBuffers.of(bytes, true, false)); } } diff --git a/tensorflow-examples/src/main/java/org/tensorflow/model/examples/dense/SimpleMnist.java b/tensorflow-examples/src/main/java/org/tensorflow/model/examples/dense/SimpleMnist.java index 994aae3..7f05dd3 100644 --- a/tensorflow-examples/src/main/java/org/tensorflow/model/examples/dense/SimpleMnist.java +++ b/tensorflow-examples/src/main/java/org/tensorflow/model/examples/dense/SimpleMnist.java @@ -56,17 +56,17 @@ public void run() { Ops tf = Ops.create(graph); // Create placeholders and variables, which should fit batches of an unknown number of images - Placeholder images = tf.placeholder(TFloat32.DTYPE); - Placeholder labels = tf.placeholder(TFloat32.DTYPE); + Placeholder images = tf.placeholder(TFloat32.class); + Placeholder labels = tf.placeholder(TFloat32.class); // Create weights with an initial value of 0 Shape weightShape = Shape.of(dataset.imageSize(), MnistDataset.NUM_CLASSES); - Variable weights = tf.variable(weightShape, TFloat32.DTYPE); + Variable weights = tf.variable(weightShape, TFloat32.class); tf.initAdd(tf.assign(weights, tf.zerosLike(weights))); // Create biases with an initial value of 0 Shape biasShape = Shape.of(MnistDataset.NUM_CLASSES); - Variable biases = tf.variable(biasShape, TFloat32.DTYPE); + Variable biases = tf.variable(biasShape, TFloat32.class); tf.initAdd(tf.assign(biases, tf.zerosLike(biases))); // Register all variable initializers for single execution @@ -98,7 +98,7 @@ public void run() { // Compute the accuracy of the model Operand predicted = tf.math.argMax(softmax, tf.constant(1)); Operand expected = tf.math.argMax(labels, tf.constant(1)); - Operand accuracy = tf.math.mean(tf.dtypes.cast(tf.math.equal(predicted, expected), TFloat32.DTYPE), tf.array(0)); + Operand accuracy = tf.math.mean(tf.dtypes.cast(tf.math.equal(predicted, expected), TFloat32.class), tf.array(0)); // Run the graph try (Session session = new Session(graph)) { @@ -108,8 +108,8 @@ public void run() { // Train the model for (ImageBatch trainingBatch : dataset.trainingBatches(TRAINING_BATCH_SIZE)) { - try (Tensor batchImages = preprocessImages(trainingBatch.images()); - Tensor batchLabels = preprocessLabels(trainingBatch.labels())) { + try (TFloat32 batchImages = preprocessImages(trainingBatch.images()); + TFloat32 batchLabels = preprocessLabels(trainingBatch.labels())) { session.runner() .addTarget(minimize) .feed(images.asOutput(), batchImages) @@ -120,16 +120,15 @@ public void run() { // Test the model ImageBatch testBatch = dataset.testBatch(); - try (Tensor testImages = preprocessImages(testBatch.images()); - Tensor testLabels = preprocessLabels(testBatch.labels()); - Tensor accuracyValue = session.runner() + try (TFloat32 testImages = preprocessImages(testBatch.images()); + TFloat32 testLabels = preprocessLabels(testBatch.labels()); + TFloat32 accuracyValue = (TFloat32)session.runner() .fetch(accuracy) .feed(images.asOutput(), testImages) .feed(labels.asOutput(), testLabels) .run() - .get(0) - .expect(TFloat32.DTYPE)) { - System.out.println("Accuracy: " + accuracyValue.data().getFloat()); + .get(0)) { + System.out.println("Accuracy: " + accuracyValue.getFloat()); } } } @@ -138,21 +137,21 @@ public void run() { private static final int TRAINING_BATCH_SIZE = 100; private static final float LEARNING_RATE = 0.2f; - private static Tensor preprocessImages(ByteNdArray rawImages) { + private static TFloat32 preprocessImages(ByteNdArray rawImages) { Ops tf = Ops.create(); // Flatten images in a single dimension and normalize their pixels as floats. long imageSize = rawImages.get(0).shape().size(); return tf.math.div( tf.reshape( - tf.dtypes.cast(tf.constant(rawImages), TFloat32.DTYPE), + tf.dtypes.cast(tf.constant(rawImages), TFloat32.class), tf.array(-1L, imageSize) ), tf.constant(255.0f) ).asTensor(); } - private static Tensor preprocessLabels(ByteNdArray rawLabels) { + private static TFloat32 preprocessLabels(ByteNdArray rawLabels) { Ops tf = Ops.create(); // Map labels to one hot vectors where only the expected predictions as a value of 1.0 diff --git a/tensorflow-examples/src/main/java/org/tensorflow/model/examples/regression/linear/LinearRegressionExample.java b/tensorflow-examples/src/main/java/org/tensorflow/model/examples/regression/linear/LinearRegressionExample.java index b8f7f24..ccd0a76 100644 --- a/tensorflow-examples/src/main/java/org/tensorflow/model/examples/regression/linear/LinearRegressionExample.java +++ b/tensorflow-examples/src/main/java/org/tensorflow/model/examples/regression/linear/LinearRegressionExample.java @@ -69,8 +69,8 @@ public static void main(String[] args) { Ops tf = Ops.create(graph); // Define placeholders - Placeholder xData = tf.placeholder(TFloat32.DTYPE, Placeholder.shape(Shape.scalar())); - Placeholder yData = tf.placeholder(TFloat32.DTYPE, Placeholder.shape(Shape.scalar())); + Placeholder xData = tf.placeholder(TFloat32.class, Placeholder.shape(Shape.scalar())); + Placeholder yData = tf.placeholder(TFloat32.class, Placeholder.shape(Shape.scalar())); // Define variables Variable weight = tf.withName(WEIGHT_VARIABLE_NAME).variable(tf.constant(1f)); @@ -97,8 +97,8 @@ public static void main(String[] args) { float y = yValues[i]; float x = xValues[i]; - try (Tensor xTensor = TFloat32.scalarOf(x); - Tensor yTensor = TFloat32.scalarOf(y)) { + try (TFloat32 xTensor = TFloat32.scalarOf(x); + TFloat32 yTensor = TFloat32.scalarOf(y)) { session.runner() .addTarget(minimize) @@ -112,31 +112,31 @@ public static void main(String[] args) { } // Extract linear regression model weight and bias values - List> tensorList = session.runner() + List tensorList = session.runner() .fetch(WEIGHT_VARIABLE_NAME) .fetch(BIAS_VARIABLE_NAME) .run(); - try (Tensor weightValue = tensorList.get(0).expect(TFloat32.DTYPE); - Tensor biasValue = tensorList.get(1).expect(TFloat32.DTYPE)) { + try (TFloat32 weightValue = (TFloat32)tensorList.get(0); + TFloat32 biasValue = (TFloat32)tensorList.get(1)) { - System.out.println("Weight is " + weightValue.data().getFloat()); - System.out.println("Bias is " + biasValue.data().getFloat()); + System.out.println("Weight is " + weightValue.getFloat()); + System.out.println("Bias is " + biasValue.getFloat()); } // Let's predict y for x = 10f float x = 10f; float predictedY = 0f; - try (Tensor xTensor = TFloat32.scalarOf(x); - Tensor yTensor = TFloat32.scalarOf(predictedY); - Tensor yPredictedTensor = session.runner() + try (TFloat32 xTensor = TFloat32.scalarOf(x); + TFloat32 yTensor = TFloat32.scalarOf(predictedY); + TFloat32 yPredictedTensor = (TFloat32)session.runner() .feed(xData.asOutput(), xTensor) .feed(yData.asOutput(), yTensor) .fetch(yPredicted) - .run().get(0).expect(TFloat32.DTYPE)) { + .run().get(0)) { - predictedY = yPredictedTensor.data().getFloat(); + predictedY = yPredictedTensor.getFloat(); System.out.println("Predicted value: " + predictedY); } diff --git a/tensorflow-examples/src/main/java/org/tensorflow/model/examples/tensors/TensorCreation.java b/tensorflow-examples/src/main/java/org/tensorflow/model/examples/tensors/TensorCreation.java index 8533d90..4041b07 100644 --- a/tensorflow-examples/src/main/java/org/tensorflow/model/examples/tensors/TensorCreation.java +++ b/tensorflow-examples/src/main/java/org/tensorflow/model/examples/tensors/TensorCreation.java @@ -28,9 +28,10 @@ * Creates a few tensors of ranks: 0, 1, 2, 3. */ public class TensorCreation { + public static void main(String[] args) { // Rank 0 Tensor - Tensor rank0Tensor = TInt32.scalarOf(42); + TInt32 rank0Tensor = TInt32.scalarOf(42); System.out.println("---- Scalar tensor ---------"); @@ -40,10 +41,10 @@ public static void main(String[] args) { System.out.println("Shape: " + Arrays.toString(rank0Tensor.shape().asArray())); - rank0Tensor.data().scalars().forEach(value -> System.out.println("Value: " + value.getObject())); + rank0Tensor.scalars().forEach(value -> System.out.println("Value: " + value.getObject())); // Rank 1 Tensor - Tensor rank1Tensor = TInt32.vectorOf(1, 2, 3, 4, 5, 6, 7, 8, 9, 10); + TInt32 rank1Tensor = TInt32.vectorOf(1, 2, 3, 4, 5, 6, 7, 8, 9, 10); System.out.println("---- Vector tensor ---------"); @@ -53,7 +54,7 @@ public static void main(String[] args) { System.out.println("Shape: " + Arrays.toString(rank1Tensor.shape().asArray())); - System.out.println("6th element: " + rank1Tensor.data().getInt(5)); + System.out.println("6th element: " + rank1Tensor.getInt(5)); // Rank 2 Tensor // 3x2 matrix of ints. @@ -63,7 +64,7 @@ public static void main(String[] args) { .set(NdArrays.vectorOf(3, 4), 1) .set(NdArrays.vectorOf(5, 6), 2); - Tensor rank2Tensor = TInt32.tensorOf(matrix2d); + TInt32 rank2Tensor = TInt32.tensorOf(matrix2d); System.out.println("---- Matrix tensor ---------"); @@ -73,7 +74,7 @@ public static void main(String[] args) { System.out.println("Shape: " + Arrays.toString(rank2Tensor.shape().asArray())); - System.out.println("6th element: " + rank2Tensor.data().getInt(2, 1)); + System.out.println("6th element: " + rank2Tensor.getInt(2, 1)); // Rank 3 Tensor // 3*2*4 matrix of ints. @@ -85,7 +86,7 @@ public static void main(String[] args) { .set(NdArrays.vectorOf(5, 6, 7, 8), 1); }); - Tensor rank3Tensor = TInt32.tensorOf(matrix3d); + TInt32 rank3Tensor = TInt32.tensorOf(matrix3d); System.out.println("---- Matrix tensor ---------"); @@ -95,6 +96,6 @@ public static void main(String[] args) { System.out.println("Shape: " + Arrays.toString(rank3Tensor.shape().asArray())); - System.out.println("n-th element: " + rank3Tensor.data().getInt(2, 1, 3)); + System.out.println("n-th element: " + rank3Tensor.getInt(2, 1, 3)); } }