jax-ml/jax

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

Code / MessageTypeSeverityTags
Malformed pickled PyTreeDef, expected 2-tuple
exception error jax, pytree, pickle
put_along_axis argument 'values' must be broadcastable to…
exception error jax, put-along-axis, broadcasting, shape-mismatch
reduce_precision: exponent_bits must be positive; got
validation error jax, reduce-precision, bit-manipulation, validation
dim of size does not match scale dim size .
exception error jax, scaled-dot, shape-validation
top_k returns int32 indices, which will overflow for array…
validation error jax, top-k, int32-overflow, large-arrays
Unsupported barrier type
exception error pallas, mosaic-gpu, interpret-mode, barrier, unsupported
Unsupported QR decomposition mode
validation error jax, linalg, qr-decomposition, argument-validation
val must be scalar.
validation error jax, pallas, triton, atomic-cas, shape
Wrong number of strides for spatial dimensions
validation error jax, numpy-reference, strides, shape-validation, convolution
axis is out of range.
exception error jax, pallas, mosaic, roll, axis-validation, shape-validation
broadcast_in_dim broadcast_dimensions must be a subset of…
validation error jax, broadcast-in-dim, index-out-of-range
broadcast_in_dim operand dimension sizes must either be 1…
validation error jax, broadcast-in-dim, shape-mismatch
full_matrices and subset_by_index cannot be both be set.
exception error jax, tpu, svd, linalg, mutually-exclusive-args
must be a tuple of factors
exception error jax, sharding, type-error, tuple
Only the POLAR (which is also DEFAULT on TPU) SVD algorithm…
exception error jax, tpu, svd, mlir-lowering, not-implemented
reduced cannot contain None. All elements in reduced should…
validation error jax, sharding, partition-spec, mesh
reduction axes contains out-of-bounds indices for .
validation error jax, lax, reduction, out-of-bounds, axis
SubViewOp only supports a single tile transform.
exception error jax, mosaic-gpu, memref, subview, tile-transform
type of weights must match type of x. Got typeof(x)=
exception error jax, bincount, weights, sharding, shape-mismatch
Unsatisfiable explicit constraint
exception error jax, shape-polymorphism, constraints, unsatisfiable
`vjp_from_jvp` is a pair of rules, not a single rule…
validation error jax, api-misuse, unpacking
Arguments to sort must have equal shapes, got
validation error jax, lax, sort, shape-mismatch
attempt to get argmax of an empty sequence
exception error jax, argmax, empty-array
Can't make a multicast reference into a peer reference.
exception error jax, pallas, mosaic-gpu, peer-memory, multicast
convolution dimension_numbers list/tuple must be length 3…
validation error jax, convolution, dimension-numbers, api-misuse
convolution_matrix: a must be at least 1-dimensional, got a…
exception error jax, convolution-matrix, input-validation, shape-error
D address calculation for multiple M tiles
exception error gpu, mosaic, tcgen05, tmem, not-implemented, tiling
dtype argument to `wald` must be a float dtype, got
validation error jax, random, dtype-validation
Expected B scales to have a M=64 collective layout, got
validation error gpu, mosaic, tcgen05, layout, collective, cgmma
invalid memory space
exception error jax, tpu, memory-space, enum, invalid-value
contains duplicated factors
validation error jax, sharding, duplicate-values, validation
`next_power_of_2` requires a non-negative integer.
validation error jax, pallas, utils, validation
num_processes must be a positive int. Got num_processes=
validation error jax, distributed, type-error, world-size
partitions cannot overlap with unreduced axes passed to…
validation error jax, sharding, partition-spec, mesh
period must be a scalar; got
exception error jax, interp, period, shape-validation
psend is currently only implemented on GPU
exception error jax, psend, backend, not-implemented
reduce_window jvp does not support non-zero…
validation error jax, autodiff, jvp, reduce-window
The 'out' argument to jnp.ptp is not supported.
validation error jax, numpy, ptp, out-argument
and must have same length.
validation error jax, pallas, mosaic-gpu, nd-loop, tiling, shape-mismatch
Unsupported data dtype
validation error jax, gpu, sparse, dtype
`vjp_from_lin` is a pair of rules, not a single rule…
validation error jax, api-misuse, unpacking
A custom return op must terminate the block.
validation error mosaic, gpu, custom-primitive, mlir, terminator
a_scale must be a TMEM Ref
validation error jax, pallas, tcgen05, block-scaling, tmem
must be a tuple or a str. Got
exception error jax, pcast, type-error, axis-name
b_scale must be a TMEM Ref
validation error jax, pallas, tcgen05, block-scaling, tmem
Because JAX arrays are immutable, jnp.ufunc.at() cannot…
exception error jax, ufunc, scatter-add, immutable-arrays
Buffer.__array__ with copy=True is not supported.
validation error jax, ffi, numpy, copy
Constraint parsing error: must contain one of '==' or '>='…
exception error jax, shape-polymorphism, constraints, parsing
Custom JVP rule must produce primal and tangent outputs…
validation error jax, custom-jvp, tangent, pytree, autodiff
dtype not understood
validation error jax, dtype, type-check, canonicalization
Effects not supported in `scan
exception error jax, scan, effects, scan3, experimental
integrate_box_1d() only handles 1D pdfs
exception error jax, scipy, kde, dimensionality
invalid distribution for variance scaling initializer
exception error jax, nn, initializers, enum-argument
make_mpi_collectives is not implemented for Windows
exception error jax, mpi, collectives, platform-support
No axis names are available. Make sure you are using…
exception error jax, pallas, mesh, collective, axis-name
size must be positive and not greater than the size of the…
exception error jax, size-validation, out-of-range
stride must be non-negative.
exception error jax, pallas, mosaic, roll, argument-validation, stride
TMEM stores expect a FragmentedArray, got
exception error mosaic, gpu, tcgen05, type-check, tensor-memory, jax
A scale layout is not supported
validation error gpu, mosaic, tcgen05, layout, tmem
at least one array or dtype is required
validation error jax, result-type, argument-validation
Can only store scalars or vectors
exception error jax, pallas, mosaic-gpu, type-error, warpgroup
Can't bitcast to
validation error mosaic-gpu, bitcast, vector-types
'devices' argument to pmap must be non-empty, or None.
validation error jax, pmap, devices, empty-argument
f32 only supports add atomics, got
exception error jax, mosaic, gpu, atomics, f32, not-implemented
from_dlpack can only unpack a dlpack tensor onto a singular…
validation error jax, dlpack, sharding, multi-device
jax.numpy.nanquantile does not support overwrite_input=True…
validation error jax, numpy, nanquantile, out-argument, overwrite-input
jax.scipy.ndimage.map_coordinates does not yet support mode
exception error jax, scipy, ndimage, unsupported-feature, boundary-mode
make_gloo_tcp_collectives only implemented for linux and…
exception error jax, gloo, collectives, platform-support
No VJP is available
validation error jax, export, vjp, serialization
Only the TMA implementation supports collective copies
exception error jax, pallas, mosaic-gpu, gpu-architecture, collective-axes, cp-async
The 'out' argument to jnp.
validation error jax, reductions, out-parameter, immutable-arrays
Unknown action
validation error jax, pallas, tpu, async-copy, internal-api
Unrecognized mode: .
validation error jax, fp8, quantization, config, cudnn
unsupported keyword arguments for mode
exception error jnp-pad, unsupported-kwarg, keyword-arguments
Vector clock size must be at least 1, but got
validation error jax, pallas, tpu, config-validation, race-detection
Acc ref dtype must be float32 or int32, got
validation error jax, pallas, tpu, accumulator, dtype
Cannot select
validation error jax, select, type-error
DLPack is only supported for devices addressable by the…
exception error jax, dlpack, device, multi-process
Each element of CompoundFactor must be a str, but got
exception error jax, shardy, sharding-rule, typeerror
Expected an accumulator ref, got
exception error jax, pallas, tpu, matmul, memory-space
expected E to be a (batched) square matrix, got E.shape=
validation error jax, linalg, frechet-derivative, shape-mismatch
MMA with element type
validation error jax, mosaic, gpu, accumulator-dtype, dtype, mma, tcgen05
Only power-of-2 num parts supported.
exception error jax, triton, pallas, split, not-implemented
Only the TMA implementation supports leader_tracked copies
exception error jax, pallas, mosaic-gpu, gpu-architecture, leader-tracked, cp-async
other cannot be a block if pointer is not a block
exception error jax, triton, pallas, load, shape, validation
`pad_width` must be of integral type.
exception error jnp-pad, type-error, integral-type
ragged_all_to_all input_offsets must be integer type.
exception error jax, ragged-all-to-all, dtype
scaled_matmul requires scales to have matching batch (B)…
validation error jax, nn, matmul, float8, scales, shape-validation
tcgen05_commit_arrive only allows arriving on a Barrier
exception error pallas, mosaic-gpu, tcgen05, barrier, interpret-mode
tpu_custom_call does not support non-trivial batching.
exception error jax, tpu, vmap, custom-call, not-implemented
Unsupported implementation option
validation error jax, nn, attention, enum-argument
A sparse metadata address calculation for multiple tiles
validation error gpu, mosaic, tcgen05, sparsity, not-implemented, tiling
correlate2d() only supports boundary='fill', fillvalue=0
exception error jax, scipy, correlation, not-implemented, boundary
Default value must be of type int, got
exception error jax, config, typeerror, int
dot_general requires rhs batch dimensions to be disjoint…
validation error jax, dot-general, dimension-numbers, batch-dims, contraction
dtype argument to jnp.std must be inexact; got
validation error jax, numpy, std, dtype-validation
group_offset is not currently supported in the…
exception error jax, pallas, gpu, not-implemented, ragged-dot
memref.StoreOp does not support transforms
exception error jax, mosaic-gpu, memref-store, transforms, not-implemented
Number of cores or threads must be at least 1, but got
validation error jax, pallas, tpu, config-validation, interpret-mode
Only one profiler server can be active at a time.
validation error jax, profiler, singleton