Module Nx_backend
include Nx_core.Backend_intf.S
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
val view : ('a, 'b) t -> Nx_core.View.tview t returns the strided view metadata describing t's logical layout (shape, strides, offset) over its underlying buffer.
val dtype : ('a, 'b) t -> ('a, 'b) Nx_core.Dtype.tdtype t returns the element type of t.
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
val buffer : context -> ('a, 'b) Nx_core.Dtype.t -> int array -> ('a, 'b) tbuffer 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.
val full : context -> ('a, 'b) Nx_core.Dtype.t -> int array -> 'a -> ('a, 'b) tfull 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, Nx_core.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, Nx_core.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, Nx_core.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, Nx_core.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, Nx_core.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, Nx_core.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, Nx_core.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, Nx_core.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
val cast : dtype:('c, 'd) Nx_core.Dtype.t -> ('a, 'b) t -> ('c, 'd) tcast ~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, Nx_core.Dtype.int32_elt) t ->
(int32, Nx_core.Dtype.int32_elt) t ->
(int32, Nx_core.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, Nx_core.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, Nx_core.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) Nx_core.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) Nx_core.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, Nx_core.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, Nx_core.Dtype.complex64_elt) t
* (Stdlib.Complex.t, Nx_core.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, Nx_core.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.
val create_context : unit -> context