jax-ml/jax

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

Code / MessageTypeSeverityTags
conv_general_dilated rhs output feature dimension size must…
validation error jax, convolution, shape-validation, grouped-conv
`device_id_type` must be MESH if `device_id` is a dict, got
exception error pallas, device-id, validation, jax
dims and idxs must have the same length
exception error jax, pallas, mosaic-gpu, cluster, validation
Expected WGStridedFragLayout, got
validation error jax, mosaic-gpu, validation, layout-mismatch, wgmma
jax.numpy.meshgrid only supports copy=True
validation error jax, meshgrid, copy-semantics, valueerror
jax.scipy.ndimage.map_coordinates currently requires…
exception error jax, scipy, ndimage, interpolation, unsupported-feature
lax.fori_loop: body_fun argument should be callable.
validation error jax, fori-loop, typeerror, argument-validation
Loading from a remote ref is only supported in jaxlib…
exception error mosaic-gpu, pallas, jaxlib-version, multi-gpu, peer-id
Mismatched number of outputs from callback. Expected
exception error jax, callback, runtime-validation, shape-mismatch
mode must be 'right' or 'left', got
validation error jax, linalg, qr-decomposition, argument-validation
can only accept axis_name which corresponds to one of…
exception error jax, pcast, named-axes, mixed-state
's input cannot be varying across the axis_name provided…
exception error jax, sharding, collectives, varying, mesh
Slicing batch dimensions is not supported.
validation error jax, pallas, indexing, slicing, not-implemented
Tiling without swizzle is not supported.
exception error
Transfer of bits is not divisible by
exception error jax, mosaic, gpu, smem, alignment, tma
Unsupported input dtype
exception error tcgen05, matmul, mosaic, gpu, dtype
Unsupported trace scope
validation error profiling, mosaic, gpu, argument-validation
Arguments to batch_matmul must be at least 2D, got
validation error jax, lax, batch-matmul, dimensionality, value-error
array() got multiple values for argument
exception error jax, argument-validation, numpy-compat, deprecation
Cannot partition over multiple dynamic parallel dimensions
exception error jax, pallas, grid-partition, dynamic-shape
Cannot transpose to
validation error jax, internal, type-dispatch
Custom JVP rule for function must produce a pair (list or…
validation error jax, custom-jvp, autodiff, return-shape
dimensions outside range
validation error jax, squeeze, axis-out-of-bounds
Expected -dim 'key' tensor for MQA. Instead got a -dim one.
validation error jax, pallas, tpu, splash-attention, shape, rank
fill_value tuple must have length equal to number of axes
validation error jax, nonzero, fill-value, argument-validation
Invalid memory space
validation error jax, pallas, tpu, memory-space, validation
mode must be one of 'full', 'valid', 'same'; got
exception error jax, convolution-matrix, argument-validation, invalid-enum-argument
next_fetch is None
exception error jax, pallas, internal-invariant, prefetch
No axis names are available. Make sure you are using…
exception error jax, pallas, mosaic-gpu, cluster, mesh, core-map
Number of device ids must match the number of mesh axes…
exception error pallas, device-id, mesh, validation, jax
Only slicing with static indices allowed
validation error mosaic, gpu, slicing, static-indices
Partitioned loads only supported for clusters of size 2…
exception error jax, pallas, mosaic-gpu, cluster-size, leader-tracked, not-implemented
ref must be a reference
exception error jax, pallas, mosaic-gpu, peer-memory, type-error
Symbolic dimension ' ' used in a context that requires a…
exception error jax, shape-polymorphism, int-conversion, tracing
The extension length n
exception error jax, scipy, stft, signal-extension, shape-validation
The shape of the accumulator
exception error jax, pallas, tpu, matmul, shape-mismatch
top_level_all_gather maintains `top_level_all_gather(x…
validation error jax, all-gather, shard-map, partition-spec
unsafe_buffer_pointer() is supported only for unsharded…
exception error jax, cuda, raw-pointer, sharding, unsupported-operation
Unsupported dot precision
exception error jax, pallas, tpu, precision, dot-general
y has more than 2 dimensions
exception error jax, statistics, covariance, shape-validation
Argument to Hessenberg reduction must have shape [..., n…
validation error jax, linalg, hessenberg, shape-validation
bfloat16 support not implemented for LSTM
exception error jax, lstm, rnn, bfloat16, cudnn, gpu
dynamic_slice: only unit steps supported in slice. Got
validation error jax, indexing, slice, stride, dynamic-slice
Got empty index range in select_range.
validation error jax, linalg, tridiagonal, argument-validation
must be positive
validation error profiling, mosaic, gpu, argument-validation
np.delete(arr, obj): got obj.dtype=
exception error jax, dtype, indexing
shard_map in_specs argument must be a pytree of…
exception error shard-map, partition-spec, none-default, jax, typeerror
Specified input which requires a copy since the source data…
validation error jax, dlpack, alignment, zero-copy
Unsupported layout for conversion from MLIR attribute
exception error jax, mosaic, layout, mlir-attribute, not-implemented
Batching over custom allocations is not supported yet.
exception error jax, pallas, vmap, batching, allocations, not-implemented
Called with a float0 array. float0s do not support any…
validation error jax, float0, autodiff, gradient, type-error
Cannot resolve attribute
exception error jax, multiref, attribute-mismatch
Commuting a `UntilingTransform` with a `ReshapeTransform`…
exception error jax, pallas, mosaic-gpu, tiling, transforms
Could not find type: .
exception error jax, pytree, pickle, registry
fiedler_companion requires the last axis of 'a' to have…
exception error jax, fiedler-companion, polynomial, empty-array
Incompatible FragmentedArray layouts
exception error jax, mosaic-gpu, layout-mismatch, pointwise, fragmented-array
index_map returned a value of type
validation error jax, pallas, index-map, block-spec
IO effect not supported in vmap-of-cond.
exception error jax, vmap, cond, io-callback, host-callback
jax.numpy.block does not allow tuples, got
validation error jax, block, tuple-vs-list, valueerror
LHS Ref must be collective if collective_axis is set.
validation error jax, pallas, tcgen05, collective-mma, tmem
Mesh has subcores, but the current TPU chip has only…
validation error jax, pallas, tpu, hardware-mismatch, sparsecore
number of KV heads must be even when megacore_mode is…
validation error jax, pallas, tpu, paged-attention, megacore, gqa
pbroadcast batcher only supports a single axis
exception error jax, pbroadcast, vmap, not-implemented
pinned array ref only works inside of a `jit`.
exception error jax, refs, pinned-memory, jit, not-implemented
Scalar arguments not (yet) supported on GPU
exception error pallas, mosaic-gpu, interpret-mode, scalar-args, not-implemented
tensordot requires axes lists to have equal length, got
validation error jax, tensordot, axes-mismatch, numpy
vmap spmd_axis_name cannot appear in shard_map out_specs
validation error jax, vmap, shard-map, spmd, out-specs
conv_general_dilated rhs output feature dimension size must…
validation error jax, lax, convolution, feature-group-count, kernel-shape
Couldn't get local_hardware_id for __dlpack__
exception error jax, dlpack, gpu, device-init, environment
Custom Partitioning rules must return Sharding.
exception error jax, custom-partitioning, sharding, tpu
` ` was passed to`canonicalize_or_get_default_platform`…
exception error jax, device-validation, type-validation
Expected exactly one collective axis, got
exception error jax, pallas, mosaic-gpu, collective-axes, leader-tracked
integer argument required; got dtype=
validation error jax, reductions, dtype, internal
jax.numpy.arange: arguments must be scalars; got
validation error jax, arange, scalar-requirement, valueerror
Lite chips, single core chips, and dual-core chips that do…
exception error tpu, chip-config, megacore, validation, jax
matmul with object of shape
exception error jax, sparse, csr, matmul, shape-validation
ormqr with left=True expects c to have the same number of…
validation error jax, linalg, ormqr, shape-validation, qr
shardings specs rank should be 3, but got lhs
validation error jax, fp8, sharding, rank-mismatch, matmul
the platform for the specified backend
exception error jax, platform-mismatch, lowering, backend
Unsigned integer dtype
exception error jax, pallas, tpu, convolution, unsigned-integer, dtype
vary_unreduced_cast is a Varying->Unreduced collective…
exception error jax, axis-name, named-axes
Weights cannot be complex types.
validation error jax, numpy, quantile, weights, complex-numbers
z.dtype= is not supported, see docstring for supported…
exception error jax, scipy-special, legendre, dtype-validation
block_size ( ) and tile_size ( ) must have the same length.
exception error jax, pallas, tpu, random, shape-mismatch
Blocked version is not implemented yet.
exception warning jax, sqrtm, not-implemented, linalg
Cannot deserialize DisabledSafetyCheck with unknown kind
exception error jax, export, deserialization, version-mismatch
Failed to infer the output layout of the iota. Please apply…
exception error pallas, mosaic-gpu, iota, layout
got , expected to be between 0 and
validation error jax, sparse, bcoo, random, nse
jnp.poch does not support complex-valued inputs.
exception error jax, scipy, poch, complex-dtype, unsupported-operation
{''.join(msg)[:-2]}
validation error jax, vmap, shape, batch-size
is not supported
exception error jax, scipy, sem, nan-handling
sources and destinations must be unique, got .
exception error jax, ppermute, collectives, validation
reduced in _specs can only be used when the mesh passed to…
exception error shard-map, partition-spec, mesh, jax, reduced
Replicated dimensions are not supported
exception error jax, mosaic, gpu, layout, replicated, warp, not-implemented
smap in_axes must be an int, None, jax.sharding.Infer, or…
exception error jax, smap, in-axes, pytree, validation
static_slice: unrecognized index
validation error jax, internal, invariant
Subclasses should implement this method
exception error jax, sharding, not-implemented, memory-kind
The on_device_size_in_bytes() method was called on
exception error jax, concretization, tracer, memory, profiling
TMEM refs must have at least 32 rows, got
exception error jax, mosaic, tmem, shape-validation, hardware-constraint
unreduced rule for is not implemented. Please file an issue…
exception error jax, sharding, gspmd, not-implemented, unreduced