Docs / nablatensor-engine-rocm / com.nablatensor.backend.rocm

final class

RocmBackend

AMD ROCm/HIP backend. Math runs in custom kernels compiled to a GCN code object at runtime with HIPRTC and launched through the HIP runtime via foreign; tensors stay resident in device memory between operations. The kernel source is GpuKernels, shared verbatim with the CUDA backend — HIPRTC accepts the CUDA C unchanged.

Phase-7 first cut: elementwise / reductions / matmul / transpose / conv2d / broadcast are implemented.

Methods

String name()
DeviceType deviceType()
boolean isAvailable()
int priority()
String deviceName()
DeviceBuffer upload(float[] data, Shape shape, DType dtype, Device device)
float[] download(DeviceBuffer buffer)
DeviceBuffer randomUniform(long seed, long counter, Shape shape, Device device)
DeviceBuffer randomNormal(long seed, long counter, Shape shape, Device device)
DeviceBuffer binary(Op op, DeviceBuffer a, DeviceBuffer b)
DeviceBuffer scalar(Op op, DeviceBuffer a, double value)
DeviceBuffer unary(Op op, DeviceBuffer a)
DeviceBuffer transpose(DeviceBuffer a)
DeviceBuffer matmul(DeviceBuffer a, DeviceBuffer b)
DeviceBuffer batchedMatmul(DeviceBuffer a, DeviceBuffer b)
DeviceBuffer reduceSum(DeviceBuffer a)
DeviceBuffer reduceMax(DeviceBuffer a)
DeviceBuffer sumAxis0(DeviceBuffer a)
DeviceBuffer sumAxis(DeviceBuffer a, int axis, boolean keepDims)
DeviceBuffer maxAxis(DeviceBuffer a, int axis, boolean keepDims)
DeviceBuffer argMaxAxis(DeviceBuffer a, int axis)
DeviceBuffer maxAxisBackward(DeviceBuffer upstream, DeviceBuffer input, int axis)
DeviceBuffer conv2d(DeviceBuffer x, DeviceBuffer w, ConvSpec spec)
DeviceBuffer conv2dGradInput(DeviceBuffer upstream, DeviceBuffer w, ConvSpec spec)
DeviceBuffer conv2dGradWeight(DeviceBuffer x, DeviceBuffer upstream, ConvSpec spec)
DeviceBuffer reluBackward(DeviceBuffer upstream, DeviceBuffer input)
DeviceBuffer fused(Expr expr, DeviceBuffer[] inputs)

Runs a whole elementwise expression chain as one generated HIP kernel launch, instead of one launch (and one intermediate device buffer) per op. The generated source is CUDA-syntax, which HIPRTC accepts unchanged; the compiled kernel is cached by its generated source, so repeated fused calls for the same op chain (any shape) reuse it without recompiling.

DeviceBuffer reshape(DeviceBuffer a, Shape target)
DeviceBuffer broadcastTo(DeviceBuffer a, Shape target)
DeviceBuffer sliceAxis0(DeviceBuffer input, int index)
DeviceBuffer stackAxis0(DeviceBuffer[] inputs)
void sync()
void release(DeviceBuffer buffer)