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));
}
}