从零实现类NumPy/ND4J多维数组库的技术咨询
Hey there! Let's break down how you can build your NumPy-like multi-dimensional array library in Java (and later Kotlin) step by step, addressing your current challenges and desired features.
1. Fixing Generic Multi-Dimensional Storage
Your initial Matrix<T> uses an Object to store values, which is neither type-safe nor efficient for numerical operations. Instead, use a flattened 1D array to hold data, paired with a shape array to track dimensions. This approach avoids Java's limitations with generic nested arrays, keeps memory contiguous (critical for performance), and simplifies index calculations.
Here's a revised base class:
import java.util.Arrays; import java.util.function.BinaryOperator; import java.util.function.Function; public class Matrix<T> { private final T[] data; private final int[] shape; private final int totalElements; @SuppressWarnings("unchecked") public Matrix(int[] shape) { this.shape = shape.clone(); // Calculate total number of elements this.totalElements = Arrays.stream(shape).reduce(1, (a, b) -> a * b); // Initialize flattened array (workaround for generic array creation) this.data = (T[]) new Object[totalElements]; } // Convert multi-dimensional indices to a flat array index private int getFlatIndex(int[] indices) { int index = 0; int stride = 1; // Calculate stride starting from the last dimension for (int i = shape.length - 1; i >= 0; i--) { index += indices[i] * stride; stride *= shape[i]; } return index; } // Get value at multi-dimensional indices public T get(int[] indices) { return data[getFlatIndex(indices)]; } // Set value at multi-dimensional indices public void set(int[] indices, T value) { data[getFlatIndex(indices)] = value; } public int[] getShape() { return shape.clone(); } }
2. Implementing Core NumPy Features
Let's add the functionality you mentioned, starting with the basics:
a. randn - Generate Normal Distributed Values
Add a static method to create a matrix filled with standard normal distribution random values (targeting Double first, since numerical operations are most common):
import java.util.Random; public static Matrix<Double> randn(int[] shape) { Matrix<Double> matrix = new Matrix<>(shape); Random rand = new Random(); for (int i = 0; i < matrix.totalElements; i++) { matrix.data[i] = rand.nextGaussian(); // Standard normal (mean 0, std 1) } return matrix; }
b. Element-Wise Operations
For operations like m*m, m+=10, or m++, implement reusable methods (Java doesn't support operator overloading, but we'll map these to Kotlin operators later):
// Element-wise multiplication public Matrix<T> multiply(Matrix<T> other, BinaryOperator<T> operation) { if (!Arrays.equals(this.shape, other.shape)) { throw new IllegalArgumentException("Shapes must match for element-wise operations"); } Matrix<T> result = new Matrix<>(this.shape); for (int i = 0; i < totalElements; i++) { result.data[i] = operation.apply(this.data[i], other.data[i]); } return result; } // Add a scalar value to all elements public Matrix<T> addScalar(T scalar, Function<T, T> addOperation) { Matrix<T> result = new Matrix<>(this.shape); for (int i = 0; i < totalElements; i++) { result.data[i] = addOperation.apply(this.data[i]); } return result; }
Use example for m*m (with Double):
Matrix<Double> m = Matrix.randn(new int[]{100, 100}); Matrix<Double> mSquared = m.multiply(m, (a, b) -> a * b);
c. Dot Product (Matrix Multiplication)
Implement matrix multiplication, ensuring dimension compatibility:
public static Matrix<Double> dot(Matrix<Double> a, Matrix<Double> b) { int[] aShape = a.getShape(); int[] bShape = b.getShape(); if (aShape[1] != bShape[0]) { throw new IllegalArgumentException( String.format("Cannot multiply matrices with shapes %s and %s", Arrays.toString(aShape), Arrays.toString(bShape)) ); } int[] resultShape = {aShape[0], bShape[1]}; Matrix<Double> result = new Matrix<>(resultShape); // Triple loop for matrix multiplication (optimize later with BLAS if needed) for (int i = 0; i < aShape[0]; i++) { for (int j = 0; j < bShape[1]; j++) { double sum = 0.0; for (int k = 0; k < aShape[1]; k++) { sum += a.get(new int[]{i, k}) * b.get(new int[]{k, j}); } result.set(new int[]{i, j}, sum); } } return result; }
3. Transition to Kotlin for Operator Overloading
Kotlin’s operator overloading makes your API feel exactly like NumPy. Here’s how to adapt the class and map methods to operators:
class Matrix<T>(val shape: IntArray) { private val data: Array<T> private val totalElements: Int init { totalElements = shape.fold(1) { acc, dim -> acc * dim } @Suppress("UNCHECKED_CAST") data = Array(totalElements) { null as T } } // Overload * for element-wise multiplication operator fun times(other: Matrix<T>, operation: (T, T) -> T): Matrix<T> { require(shape.contentEquals(other.shape)) { "Shapes must match for element-wise multiplication" } val result = Matrix<T>(shape) for (i in data.indices) { result.data[i] = operation(data[i], other.data[i]) } return result } // Overload + for scalar addition operator fun plus(scalar: T, operation: (T) -> T): Matrix<T> { val result = Matrix<T>(shape) for (i in data.indices) { result.data[i] = operation(data[i]) } return result } // Dot product (matrix multiplication) companion object { fun dot(a: Matrix<Double>, b: Matrix<Double>): Matrix<Double> { require(a.shape[1] == b.shape[0]) { "Matrix dimensions mismatch for dot product" } val resultShape = intArrayOf(a.shape[0], b.shape[1]) val result = Matrix<Double>(resultShape) for (i in 0 until a.shape[0]) { for (j in 0 until b.shape[1]) { var sum = 0.0 for (k in 0 until a.shape[1]) { sum += a.get(intArrayOf(i, k)) * b.get(intArrayOf(k, j)) } result.set(intArrayOf(i, j), sum) } } return result } } // Helper methods for get/set fun get(indices: IntArray): T { var index = 0 var stride = 1 for (i in shape.indices.reversed()) { index += indices[i] * stride stride *= shape[i] } return data[index] } fun set(indices: IntArray, value: T) { var index = 0 var stride = 1 for (i in shape.indices.reversed()) { index += indices[i] * stride stride *= shape[i] } data[index] = value } }
Now you can write NumPy-style code:
// Create a 100x100 matrix with random normal values val m = Matrix.randn(intArrayOf(100, 100)) // Element-wise multiplication: m * m val mSquared = m times m { a, b -> a * b } // Add 10 to all elements: m +=10 val mPlus10 = m plus 10.0 { it + 10.0 } // Dot product: np.dot(m, m) val dotProduct = Matrix.dot(m, m)
4. Additional Tips for Improvement
- Type-Specific Subclasses: For better performance, create specialized classes like
DoubleMatrixorIntMatrixinstead of using genericMatrix<T>. This avoids lambda overhead and lets you use direct numerical operations. - Performance Optimization: For large matrices, integrate native linear algebra libraries (like OpenBLAS) via JNI, similar to how ND4J works.
- Shape Validation: Add helper methods to validate shape compatibility for all operations (e.g., broadcasting, if you want to support it later).
- Unit Testing: Write tests to verify your implementation matches NumPy’s output for operations like dot product and
randn.
内容的提问来源于stack exchange,提问作者Clement

