jax-ml/jax

Documented errors, page 22 of 23. Back to jax-ml/jax

Code / MessageTypeSeverityTags
The layout of ShapedArray should not be `AutoLayout` when…
validation error jax, layout, auto-layout, configuration
The 'raise' mode to jnp.take is not supported.
validation error jax, take, mode, bounds-check
4-bit block scaled MMA only supports K-fastest operands…
validation error gpu, mosaic, tcgen05, layout, mxfp4, block-scaling
accumulate only supported for binary ufuncs
exception error jax, ufunc, accumulate, api-misuse
B scale shape[0] must be a multiple of 128 and >= N=
validation error gpu, mosaic, tcgen05, shape-mismatch, block-scaling, alignment
Cannot concatenate vectors of different element types
validation error mosaic-gpu, vector-concat, dtype-mismatch
dot_general requires lhs batch dimensions to be disjoint…
exception error jax, dot-general, dimension-numbers, batch-dims, contraction
Dynamic grid bounds not (yet) supported on GPU
exception error pallas, mosaic-gpu, interpret-mode, dynamic-grid, not-implemented
Expected B scales to have a M=128 layout, got
validation error gpu, mosaic, tcgen05, layout, block-scaling
expected , got
validation error jax, linearize, jvp, pytree
get_topology_for_devices requires >= 1 devices.
validation error jax, topology, devices, validation
Invalid CSR buffer sizes
validation error jax, sparse, spsolve, csr, buffer-sizes
invalid mode for variance scaling initializer
exception error jax, nn, initializers, enum-argument
JAX does not support any version below
exception error jax, dlpack, version-negotiation, protocol
memref.cast tmem layouts must be identical for both input…
exception error jax, mosaic-gpu, tmem, cast, layout
No batching rule defined for custom_vmap function
exception error jax, custom-vmap, vmap, missing-rule
Only support preferred_element_type in (f32, bf16, f16)…
validation error jax, dtype, fp8, matmul, unsupported-type
Python int too large to convert to
validation error jax, overflow, int64, x64
reduce_window got the wrong number of window_dimensions for…
validation error jax, shape-validation, reduce-window, windowing
Reshape ref with dynamic size is not supported.
exception error jax, reshape, dynamic-shape
series_order must be non-negative.
validation error jax, scipy-special, log-ndtr, argument-validation
Shape and strides must have the same length
exception error mosaic, fragmented-array, tiling, shape-strides, validation
sparse_format= not recognized; must be one of
exception error jax, sparse, argument-validation, format-string
The shape of the accumulator
exception error jax, pallas, tpu, matmul, shape-mismatch
Tiles must not be empty
validation error mosaic, fragmented-array, tiling, validation
trace_value requires a scalar value, got shape
exception error jax, pallas, tpu, debugging, scalar-required
Types must match, got
validation error mosaic-gpu, prmt, type-mismatch
array split does not result in an equal division: rest is
exception error jax, split, uneven-division
Clip received a complex value either through the input or…
exception error jax, clip, complex-dtype, unsupported-operation
expected A to be a (batched) square matrix, got A.shape=
validation error jax, linalg, frechet-derivative, shape-mismatch
jax.numpy.put_along_axis cannot modify arrays in-place…
exception error jax, put-along-axis, immutability, inplace
K mismatch: !=
exception error jax, mosaic, mma, shape-mismatch
lax.platform_dependent: the 'default' branch must be a…
validation error jax, lax, platform-dependent, typeerror, callable
logical reduction requires operand dtype bool or int, got
validation error jax, lax, logical-reduction, dtype
masked swap with strided store
exception error jax, pallas, tpu, strided-store, masked-store
multiple dimensions cannot be all_gathered since…
validation error jax, all-gather, shard-map, multi-dim
Must signal an int32 value, but got
exception error pallas, semaphore, dtype-validation, jax
precv is currently only implemented on GPU
exception error jax, precv, backend, not-implemented
ragged_all_to_all output_offsets must be integer type.
exception error jax, ragged-all-to-all, dtype
Sharding spec implies that array axis is partitioned times…
exception error jax, sharding, named-sharding, mesh, partitionspec
sparse rule for is not implemented.
exception error jax, sparse, sparsify, not-implemented, lax
The denominator cannot be unreduced passed to `div`. Got
exception error jax, sharding, div, unreduced, value-error
The 'out' argument to jnp.nanmean is not supported.
validation error jax, numpy, nanmean, out-argument
The size of all_to_all split_axis
exception error jax, all-to-all, shape, spmd
Use 'cuda', 'rocm', or 'oneapi' for lax.platform_dependent.
validation error jax, lax, platform-dependent, gpu, invalid-argument
x argument to bincount must have an integer type; got
exception error jax, bincount, dtype, typeerror
A scale address calculation for multiple M tiles
validation error gpu, mosaic, tcgen05, block-scaling, not-implemented, tiling
B scale address calculation for multiple N tiles
validation error gpu, mosaic, tcgen05, block-scaling, not-implemented, tiling
Invalid dimension for tiling
validation error mosaic, fragmented-array, tiling, index-out-of-range
Either both or neither of the x and y arguments should be…
exception error jax, where, missing-argument
Invalid mode ' ' for np.take
validation error jax, take, invalid-argument, mode
jnp.unwrap does not support complex inputs.
exception error jax, unwrap, complex-dtype, unsupported-operation
Dropout not supported in LSTM reference because we cannot…
exception error jax, lstm, rnn, dropout, not-implemented
Expected a 3-dim mask, instead got
validation error jax, splash-attention, mask, rank-mismatch, tpu
Only M=128 and M=64 are supported for MMA, but got M=
validation error jax, mosaic, gpu, mma, shape-validation, tcgen05
ragged_all_to_all recv_sizes must be integer type.
exception error jax, ragged-all-to-all, dtype
must satisfy <=start<=
exception error jax, rollaxis, argument-validation
condlist must have length equal to choicelist
exception error jax, select, length-mismatch
reduce_window got inconsistent base_dilation and…
validation error jax, shape-validation, reduce-window, dilation
Cannot concatenate non-vector values
validation error mosaic-gpu, vector-concat, type-validation
PTX does not support unsigned WGMMA accumulators
validation error jax, mosaic-gpu, wgmma, ptx, signedness
not valid, must be one of [1, 2, 3, 4]
validation error jax, sparse, spsolve, gpu, parameter-validation
Argument to get_c_api_topology was not a pjrt_c_api capsule.
validation error jax, pjrt, topology, capsule, validation
Cannot do a non-empty jnp.take() from an empty axis.
validation error jax, take, empty-array, out-of-bounds
code argument must be a code object
validation error jax, traceback, type-check
Dimension must be either 2 or 3 for cross product
exception error jax, cross-product, shape-validation
Dtype mismatch: !=
exception error jax, mosaic, mma, dtype-mismatch
k argument to top_k must be no larger than size along axis…
validation error jax, top-k, shape-validation, off-by-one
No swizzle is not supported
exception error tcgen05, matmul, swizzle, shared-memory
`strides` must contain only 1s.
exception error jax, mosaic, gpu, vector, strides
Only unstack along the last dimension is supported in…
exception error jax, triton, pallas, unstack, axis, not-implemented
correlate2d() only supports 2-dimensional inputs.
exception error jax, scipy, correlation, rank-mismatch, shape-validation
Arrays must be one-dimensional. Got
validation error jax, sparse, spsolve, shape, dimensions
explicit tiling is only supported for SparseCore kernels.
exception error tpu, pallas, tiling, sparsecore, invalid-argument
x must be a one-dimensional array
exception error jax, vander, shape-validation
Argument to get_c_api_topology contained a null pointer.
validation error jax, pjrt, topology, capsule, null-pointer
attempt to get argmin of an empty sequence
exception error jax, argmin, empty-array
WGMMA instruction only supports f32, f16 and s32 out
validation error jax, mosaic-gpu, wgmma, dtype, accumulator
Expected an input array of integer or boolean data type
exception error jax, packbits, dtype-validation
`slice_index` has been deprecated. Please use…
console warning jax, distributed, api-rename, deprecation
Only signed accumulator supported for integer operands.
exception error jax, mosaic, mma, accumulator, signedness
Unsupported A register array dtype
validation error jax, mosaic-gpu, wgmma, dtype, registers
Scale element type mismatch: expected f8e8m0fnu or…
validation error gpu, mosaic, tcgen05, dtype, block-scaling
dtype parameter is not supported by Buffer.__array__.
validation error jax, ffi, numpy, dtype
reduceat only supported for binary ufuncs
exception error jax, ufunc, reduceat, api-misuse
chunk_size must be positive
validation error jax, splash-attention, mask, chunked-attention, argument-validation
The 'out' argument to jnp.logaddexp2.reduce is not…
validation error jax, logaddexp2, out-parameter
4-bit block scaled MMA only supports K-fastest operands…
validation error gpu, mosaic, tcgen05, layout, mxfp4, block-scaling
The 'out' argument to jnp.logaddexp.reduce is not supported.
validation error jax, logsumexp, out-parameter
Unexpected mask shape, got
validation error jax, splash-attention, mask, shape-mismatch, multi-head
Malformed pickled FlattenedIndexKey, expected 1-tuple
exception error jax, pytree, pickle
scale must be None, 'sqrtn', or 'n'; got
exception error jax, scipy, linalg, dft, argument-validation
m_warps must be 1, 2, or 4, but got
exception error jax, mosaic, mma, warps, validation
top_k operand must have >= 1 dimension, got
validation error jax, top-k, scalar, rank-error
Only 32-bit scalar types supported
validation error mosaic-gpu, redux, dtype, hardware-limit
The input must be non-scalar to take a cumulative product…
validation error jax, numpy, cumulative-prod, scalar-input
explicit opt_level is only supported for SparseCore kernels.
exception error tpu, pallas, opt-level, sparsecore, invalid-argument
The input must be non-scalar to take a cumulative sum…
validation error jax, numpy, cumulative-sum, scalar-input
MMA with element type
validation error jax, mosaic, gpu, accumulator-dtype, dtype, mma, tcgen05
'order' must be either 'little' or 'big'
exception error jax, packbits, bitorder, argument-validation