Module type Backend_intf.S
Backend interface for Nx tensor operations.
This module type defines the contract between Nx's frontend and its pluggable backends. Backends may execute operations eagerly (C backend), raise effects for JIT compilation (Rune), build computation graphs, or implement other execution strategies.
Design Philosophy
Operations exist at the level of C standard library functions: every operation that maps to a C stdlib call is a backend primitive, avoiding the overhead of composing multiple operations in eager mode. Rune's JIT pipeline can decompose these into lower primitives when building computation graphs.
Frontend/Backend Contract
The frontend is responsible for:
- Broadcasting inputs to matching shapes before calling binary operations.
- Promoting dtypes to compatible types before calling operations.
- Validating parameters (axes in range, shapes compatible, etc.).
The backend can assume all inputs are well-formed. It is responsible for:
- Executing the operation correctly for all supported dtypes.
- Handling strided (non-contiguous) inputs via the view metadata.
- Returning tensors with correct view metadata.
Conventions
- All compute operations allocate and return their result. The frontend passes pre-broadcasted, pre-validated inputs and receives the result tensor.
- Movement operations manipulate view metadata (shape, strides, offset) without copying data when possible.
Types
'a is the OCaml element type (e.g., float, int32). 'b is a phantom type that tags the dtype for type safety.
Backend execution context.
Carries backend-specific state such as memory pools, device handles, command queues, or computation graphs.
Tensor Properties
view t returns the strided view metadata describing t's logical layout (shape, strides, offset) over its underlying buffer.
val to_host : ('a, 'b) t -> ('a, 'b) Nx_buffer.tto_host t returns t's data as a flat, C-contiguous host buffer.
Use view to interpret the logical structure. CPU backends may return a direct reference (zero-copy); GPU backends copy from device to host.
Tensor Creation
buffer ctx dtype shape allocates an uninitialized tensor.
Contents are undefined. Used internally by backends to allocate output tensors.
Backend must: return a tensor with the given shape and dtype whose view is C-contiguous.
full ctx dtype shape value creates a tensor where every element is value.
For scalars, shape is [||]. Subsumes zeros, ones, and constant fill.
Backend must: return a C-contiguous tensor of the given shape and dtype with all elements set to value.
val from_host : context -> ('a, 'b) Nx_buffer.t -> ('a, 'b) tfrom_host ctx buf creates a tensor from a flat, C-contiguous host buffer.
CPU backends may share the buffer directly (zero-copy). GPU backends copy from host to device.
Frontend guarantees: buf is C-contiguous.
Element-wise Binary Operations
Frontend guarantees: a and b have identical shapes (after broadcasting) and compatible dtypes (after promotion).
Backend must: allocate a C-contiguous output tensor with the correct shape and write the result.
Arithmetic
div a b is the element-wise quotient of a and b.
Integer dtypes use truncation toward zero (C division). Floating-point dtypes use IEEE 754 division.
mod_ a b is the element-wise remainder of a / b.
Integers use C's % operator (truncated division). Floats use fmod. The sign of the result follows the dividend a.
pow base exponent is the element-wise power base ^ exponent.
atan2 y x is the element-wise arc tangent of y / x.
Returns the angle in radians in (-π, π], handling all quadrants.
Comparison
Comparison operations produce boolean tensors.
val cmpeq : ('a, 'b) t -> ('a, 'b) t -> (bool, Dtype.bool_elt) tcmpeq a b is the element-wise equality test of a and b.
val cmpne : ('a, 'b) t -> ('a, 'b) t -> (bool, Dtype.bool_elt) tcmpne a b is the element-wise inequality test of a and b.
val cmplt : ('a, 'b) t -> ('a, 'b) t -> (bool, Dtype.bool_elt) tcmplt a b is the element-wise less-than test of a and b.
val cmple : ('a, 'b) t -> ('a, 'b) t -> (bool, Dtype.bool_elt) tcmple a b is the element-wise less-or-equal test of a and b.
Min/Max
Bitwise
Operate on the binary representation of integer and boolean dtypes. For booleans, these are equivalent to logical AND/OR/XOR.
and_ a b is the element-wise bitwise AND of a and b.
Element-wise Unary Operations
Frontend guarantees: x has compatible dtype.
Backend must: allocate a C-contiguous output tensor with the correct shape and write the result.
Arithmetic
sign x is the element-wise sign of x: -1 for negative, 0 for zero, 1 for positive. Returns NaN for floating-point NaN inputs.
Exponential and Logarithm
Trigonometric
All inputs are in radians.
asin x is the element-wise arc sine of x.
Returns values in [-π/2, π/2].
acos x is the element-wise arc cosine of x.
Returns values in [0, π].
atan x is the element-wise arc tangent of x.
Returns values in [-π/2, π/2].
Hyperbolic
Rounding
For integer dtypes, all rounding operations are the identity.
round x rounds each element to nearest integer, half away from zero (C's round).
Special Functions
Ternary Operations
val where : (bool, Dtype.bool_elt) t -> ('a, 'b) t -> ('a, 'b) t -> ('a, 'b) twhere cond if_true if_false selects elements: if_true.{i} where cond.{i} is true, if_false.{i} otherwise.
Frontend guarantees: all three input tensors have identical shapes. cond is boolean. if_true and if_false share the same dtype.
Reduction Operations
Reductions aggregate values along one or more axes.
Frontend guarantees: axes contains valid, non-negative, deduplicated axis indices.
reduce_sum ~axes ~keepdims x sums elements of x along axes.
reduce_prod ~axes ~keepdims x multiplies elements of x along axes.
reduce_max ~axes ~keepdims x finds the maximum of x along axes.
reduce_min ~axes ~keepdims x finds the minimum of x along axes.
val argmax :
axis:int ->
keepdims:bool ->
('a, 'b) t ->
(int32, Dtype.int32_elt) targmax ~axis ~keepdims x returns int32 indices of maximum values of x along axis. For ties, returns the first occurrence.
Frontend guarantees: axis is valid and non-negative.
val argmin :
axis:int ->
keepdims:bool ->
('a, 'b) t ->
(int32, Dtype.int32_elt) targmin ~axis ~keepdims x returns int32 indices of minimum values of x along axis. For ties, returns the first occurrence.
Frontend guarantees: axis is valid and non-negative.
associative_scan ~axis ~op x computes an inclusive prefix scan of x along axis. `Sum for cumulative sum, `Prod for cumulative product, `Max/`Min for running max/min.
Frontend guarantees: axis is valid and non-negative.
Sort Operations
Frontend guarantees: axis is valid and non-negative.
sort ~axis ~descending x sorts elements of x along axis. NaN values are placed at the end regardless of sort direction.
val argsort :
axis:int ->
descending:bool ->
('a, 'b) t ->
(int32, Dtype.int32_elt) targsort ~axis ~descending x returns int32 indices that would sort elements of x along axis.
Movement Operations
Movement operations manipulate view metadata (shape, strides, offset) without copying data when possible. They return new tensor handles sharing the underlying buffer.
Frontend guarantees: all parameters are validated (axes in range, shapes compatible, bounds within limits).
Backend must: return a tensor with the correct view metadata. May share the underlying buffer (zero-copy) or allocate if necessary.
expand t shape broadcasts dimensions of size 1 to match shape by setting their stride to 0. Non-singleton dimensions must already match. Zero-copy.
reshape t shape changes the logical shape, preserving element count.
Zero-copy when t is C-contiguous or the reshape is compatible with the current strides. May copy if t is non-contiguous.
permute t axes reorders dimensions according to axes, which must be a permutation of [0, ..., ndim-1]. Zero-copy.
shrink t ranges extracts a contiguous slice. ranges.(i) is (start, stop) with exclusive stop. Zero-copy (adjusts offset and shape).
flip t axes reverses dimensions where axes.(i) = true by negating strides. Zero-copy.
pad t padding fill_value extends t with fill_value. padding.(i) is (before, after) for dimension i.
Backend must: allocate a new buffer and copy data.
cat tensors ~axis concatenates tensors along axis.
Frontend guarantees: all tensors have the same shape except along axis. axis is valid. The list is non-empty.
Type Conversion and Memory
cast ~dtype x converts elements of x to dtype.
Float-to-int truncates toward zero. Int-to-float may lose precision for large values.
contiguous t returns a C-contiguous version of t.
May return t unchanged if already C-contiguous. Otherwise allocates and copies.
Backend must: return a C-contiguous tensor with the same data.
copy t creates an independent copy with its own buffer.
Backend must: always allocate a new buffer, even if t is already contiguous.
assign dst src copies elements from src into dst in-place.
Frontend guarantees: dst and src have matching shapes and dtypes.
Backend must: write src's data into dst's buffer, respecting both tensors' strides.
Random Number Generation
val threefry :
(int32, Dtype.int32_elt) t ->
(int32, Dtype.int32_elt) t ->
(int32, Dtype.int32_elt) tthreefry key counter applies the Threefry-2x32 hash function.
Frontend guarantees: key and counter are int32 tensors with compatible shapes.
Indexed Access Operations
val gather : ('a, 'b) t -> (int32, Dtype.int32_elt) t -> axis:int -> ('a, 'b) tgather data indices ~axis selects elements from data along axis using indices.
Frontend guarantees: rank data = rank indices. axis is valid. Index values are in range for data's size along axis.
val scatter :
?mode:[ `Set | `Add ] ->
?unique_indices:bool ->
('a, 'b) t ->
indices:(int32, Dtype.int32_elt) t ->
updates:('a, 'b) t ->
axis:int ->
('a, 'b) tscatter ?mode ?unique_indices template ~indices ~updates ~axis places updates into a tensor shaped like template along axis.
`Set (default) uses the last update for duplicate indices. `Add accumulates every update into the template's value. unique_indices = true hints that indices are unique.
Frontend guarantees: rank indices = rank updates. axis is valid. template has the desired output shape.
Backend must: allocate and return the result tensor, initialized from template's data.
Window Operations
Sliding-window extraction and its inverse. Used to implement convolution as unfold + reshape + matmul and pooling as unfold + reduce.
val unfold :
('a, 'b) t ->
kernel_size:int array ->
stride:int array ->
dilation:int array ->
padding:(int * int) array ->
('a, 'b) tunfold t ~kernel_size ~stride ~dilation ~padding extracts sliding windows from the last K spatial dimensions, where K = Array.length kernel_size.
Input shape (leading..., spatial...) produces (leading..., prod(kernel_size), L) where L is the number of windows. All dimensions before the last K are preserved as-is.
Frontend guarantees: all array parameters have length K. Values are positive. Input has at least K dimensions.
Backend must: allocate and return the result tensor.
val fold :
('a, 'b) t ->
output_size:int array ->
kernel_size:int array ->
stride:int array ->
dilation:int array ->
padding:(int * int) array ->
('a, 'b) tfold t ~output_size ~kernel_size ~stride ~dilation ~padding combines sliding windows (inverse of unfold). Overlapping values are summed.
Input shape (leading..., prod(kernel_size), L) produces (leading..., output_size...).
Frontend guarantees: parameters are consistent with a valid unfold configuration.
Backend must: allocate and return the result tensor.
Matrix Operations
matmul a b computes matrix multiplication a × b.
For 2D inputs: standard matrix multiply. For higher dimensions: batched multiply on the last two dimensions, with broadcasting via strides.
Frontend guarantees: a's last dim equals b's second-to-last dim.
Backend must: allocate and return the result. May use BLAS for performance. a and b may be non-contiguous.
Fourier Transforms
Frontend guarantees: axes contains valid, non-negative axis indices. Input tensors have compatible complex or real dtypes.
val fft :
?out:(Stdlib.Complex.t, 'b) t ->
(Stdlib.Complex.t, 'b) t ->
axes:int array ->
(Stdlib.Complex.t, 'b) tfft ?out t ~axes computes the forward DFT along axes.
val ifft :
?out:(Stdlib.Complex.t, 'b) t ->
(Stdlib.Complex.t, 'b) t ->
axes:int array ->
(Stdlib.Complex.t, 'b) tifft ?out t ~axes computes the inverse DFT along axes.
val rfft :
?out:(Stdlib.Complex.t, 'b) t ->
(float, 'a) t ->
dtype:(Stdlib.Complex.t, 'b) Dtype.t ->
axes:int array ->
(Stdlib.Complex.t, 'b) trfft ?out t ~dtype ~axes computes the real-input DFT along axes.
Exploits conjugate symmetry to return only the non-redundant half of the spectrum along the last transformed axis.
val irfft :
?out:(float, 'b) t ->
?s:int array ->
(Stdlib.Complex.t, 'a) t ->
dtype:(float, 'b) Dtype.t ->
axes:int array ->
(float, 'b) tirfft ?out ?s t ~dtype ~axes computes the inverse real-input DFT along axes.
Takes conjugate-symmetric complex input, returns real output. s specifies output sizes along the transformed axes; None infers sizes from the input.
Linear Algebra
All linalg operations support batching: the last two dimensions are the matrix dimensions, earlier dimensions are batch dimensions.
Frontend guarantees: input matrices have compatible shapes (square where required, matching dimensions for solves).
Backend must: allocate and return result tensors. Typically delegates to LAPACK.
cholesky ~upper t computes the Cholesky factorization of a positive-definite matrix. Returns L (lower) or U (upper) such that A = L·Lᵀ or A = Uᵀ·U.
qr ~reduced t returns (Q, R) where Q is orthogonal and R is upper triangular. reduced = true returns economy-size factorization.
val svd :
full_matrices:bool ->
('a, 'b) t ->
('a, 'b) t * (float, Dtype.float64_elt) t * ('a, 'b) tsvd ~full_matrices t returns (U, S, Vᴴ). S is a 1D float64 vector of singular values in descending order. full_matrices = false returns thin SVD.
val eig :
vectors:bool ->
('a, 'b) t ->
(Stdlib.Complex.t, Dtype.complex64_elt) t
* (Stdlib.Complex.t, Dtype.complex64_elt) t optioneig ~vectors t computes eigenvalues (and optionally eigenvectors) of a square matrix. Returns complex64 results.
val eigh :
vectors:bool ->
('a, 'b) t ->
(float, Dtype.float64_elt) t * ('a, 'b) t optioneigh ~vectors t computes eigenvalues (and optionally eigenvectors) of a symmetric/Hermitian matrix. Eigenvalues are float64.
val triangular_solve :
upper:bool ->
transpose:bool ->
unit_diag:bool ->
('a, 'b) t ->
('a, 'b) t ->
('a, 'b) ttriangular_solve ~upper ~transpose ~unit_diag a b solves A·x = b or Aᵀ·x = b where A is triangular.
upper: A is upper triangular. transpose: solve Aᵀ·x = b. unit_diag: assume diagonal is all ones.