jax-ml/jax
Documented errors, page 11 of 23. Back to jax-ml/jax
| Code / Message | Type | Severity | Tags |
|---|---|---|---|
| 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 |