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