You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

从零实现类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 DoubleMatrix or IntMatrix instead of using generic Matrix<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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.05.12 04:53:03