jax-ml/jax
Documented errors, page 7 of 23. Back to jax-ml/jax
| Code / Message | Type | Severity | Tags |
|---|---|---|---|
| Invalid value " " for JAX flag | validation | error | jax, config, valueerror, environment-variable, enum |
| jvp called with different primal and tangent shapes;Got… | validation | error | jax, jvp, shape, autodiff |
| mean does not have dimension | exception | error | jax, scipy, kde, shape-validation |
| `nextafter` only supports float32 and float64, but got | validation | error | jax, pallas, triton, dtype, nextafter |
| must be divisible by | validation | error | jax, pallas, tpu, gqa, attention, head-dimension |
| Partitioned callback not supported with return values. | exception | error | jax, callback, sharding, api-contract |
| Reduction op not supported by the TMA implementation for… | validation | error | jax, mosaic-gpu, tma, reduction, dtype-mismatch, unsupported-operation |
| The minor dimension size of an accumulator ref must be | validation | error | jax, pallas, tpu, accumulator, shape |
| Unsupported gather | exception | error | jax, pallas, tpu, gather, not-implemented |
| Attempting to store into allocation with key | exception | error | jax, pallas, mosaic, shared-memory, key-collision, interpret-mode |
| cannot reshape array of shape | exception | error | jax, reshape, size-mismatch |
| Cannot stack arrays with different numbers of dimensions… | validation | error | jax, stack, rank-mismatch, shape-validation |
| `collective_axes` must be specified when `leader_tracked`… | validation | error | mosaic-gpu, pallas, api-misuse, collectives |
| Equation must contain exactly one '->' | validation | error | jax, einshape, einsum-notation, equation-format, parse-error |
| for grad support, subclass | validation | error | jax, autodiff, custom-primitive, not-implemented |
| must be less or equal to | validation | error | jax, pallas, tpu, ragged-attention, capacity-validation |
| QDWH implementation is only supported on TPU | validation | error | jax, eigh, qdwh, tpu, backend, not-implemented |
| Unknown polar decomposition method | exception | error | jax, argument-validation, polar-decomposition |
| Unsupported memory space. | exception | error | jax, pallas, mosaic-gpu, memory-space, aliasing, not-implemented |
| axes argument to transpose() | exception | error | jax, sparse, coo, transpose, not-implemented |
| bcoo_slice: indices must have size mat.ndim= | exception | error | jax, sparse, bcoo, validation, shape-mismatch |
| can only convert to extended dtype from an array of its… | exception | error | jax, extended-dtype, shape-validation, suffix-mismatch |
| coordinates must be a sequence of length input.ndim, but | exception | error | jax, scipy, ndimage, coordinates, shape-mismatch |
| Currently only support batch_dim in [0, None], but got | exception | error | jax, vmap, cudnn, batching, in-axes |
| Effects not supported in AD of `checkpoint`/`remat | error_code | error | pytree, none-handling, jax, tree-map, breaking-change |
| Expected unreduced_kind to be of type… | validation | error | jax, sharding, unreduced-kind, type-error |
| `fan_in` must be less or equal than `fan_out`. | exception | error | jax, initializer, shape-validation, fan-in-out |
| cannot accept args which are unreduced. Got and axes= | validation | error | jax, sharding, collectives, unreduced, spmd |
| Unsupported ndim | validation | error | jax, pallas, ndim, shape-validation |
| When LU decomposition matrix and b different numbers of… | validation | error | jax, lu-solve, shape-validation, broadcasting |
| corrcoef: dtype must be a subclass of float or complex; got | exception | error | jax, corrcoef, dtype-validation |
| Expected input and output shapes are the same after… | exception | error | jax, bitcast, divisibility |
| Folding dimensions starting from is out of bounds for shape | validation | error | jax, mosaic-gpu, memref, shape-mismatch, index-out-of-bounds |
| name must be non-empty | validation | warning | mosaic, gpu, io, dump, argument-validation |
| Python int too large to convert to int64 | validation | error | pytree, custom-node, jax, tree-flatten-with-path |
| Argument to symmetric eigendecomposition must have shape… | validation | error | jax, linalg, eigh, shape-validation, square-matrix |
| Arguments to jax.numpy.lcm must be integers. | exception | error | jax, numpy, dtype-validation, integer-required |
| BCSR from_scipy_sparse requires 2D array; | exception | error | jax, sparse, bcsr, scipy, input-validation |
| Expected same element type, got | validation | error | jax, mosaic-gpu, async-copy, dtype-mismatch |
| Expected slice start | exception | error | jax, pallas, alignment, tiling, slicing |
| External meshes are not supported by the Mosaic GPU backend | exception | error | pallas, mosaic-gpu, mpmd-map, mesh, not-implemented |
| hessenberg requires the last dimension of a to be constant… | validation | error | jax, hessenberg, dynamic-shapes, jit, cpu |
| index out of bounds for axis with size ( ) | validation | error | jax, indexing, out-of-bounds |
| linearized function called on tangent values inconsistent… | validation | error | jax, linearize, tangent, aval, mixed-precision |
| ndim should be , but got | validation | error | jax, nn, attention, shape-validation |
| shape should be : but got | validation | error | jax, nn, attention, shape-validation |
| wrapped function must be passed at least one argument… | validation | error | jax, vmap, axis-size, api-misuse |
| pallas_call does not support hijax for index_map | exception | error | jax, pallas, pallas-call, index-map, lowering, notimplementederror |
| run_scoped lowering outside of Pallas does not support… | exception | error | pallas, run-scoped, collectives, jax |
| Async copies only support striding up to 5 dimensions | validation | error | jax, mosaic-gpu, tma, rank-limit, shape-validation |
| Batching with multiple indexers not supported. | exception | error | jax, vmap, batching, not-implemented, multiple-indexing |
| cannot reshape array of shape | exception | error | jax, reshape, size-mismatch |
| CSC.tree_unflatten: invalid | exception | error | jax, sparse, pytree, csc, serialization |
| dot_general requires lhs dimension numbers to be… | exception | error | jax, dot-general, dimension-numbers, index-out-of-range |
| Either both or neither `src_sem` and `device_id` can be set. | exception | error | jax, pallas, mosaic, dma, remote-copy, argument-validation |
| indices must have an integer type | exception | error | jax, lax, gather, dtype |
| Leading dimension of seed key_data must be 1. | validation | error | jax, pallas, tpu, prng, shape-validation |
| No cluster axes found. | validation | error | jax, pallas, tcgen05, cluster, collective-mma |
| None is not a valid value for jnp.array | exception | error | jax, null-handling, data-validation |
| the only valid string value of `left` is 'extrapolate', but… | exception | error | jax, interp, invalid-argument, sentinel-value |
| Unsupported shape: . TMEM references must have either or… | exception | error | jax, mosaic, tmem, shape-validation, blackwell |
| `weights` input should be of length n | exception | error | jax, scipy, kde, weights, length-mismatch |
| Axes mentioned in `manual_axis_type` field of ShapedArray… | exception | error | jax, sharding, manual-axis-type, mesh, validation |
| `bw_method` should be 'scott', 'silverman', a scalar, or a… | exception | error | jax, scipy, kde, bandwidth, argument-validation |
| Can only store to references | exception | error | jax, pallas, mosaic-gpu, type-error, warpgroup |
| Cannot on a non-()-shaped semaphore | exception | error | pallas, semaphore, shape-validation, jax |
| cluster= must be at most 3D, got | validation | error | jax, pallas, mosaic-gpu, cluster, launch-config, validation |
| Expected list, got . | validation | error | jax, pytree, type-mismatch, list |
| expected w and y to have the same length | validation | error | jax, numpy, polynomial, polyfit, weights, length-mismatch |
| JAX does not support string indexing; got | validation | error | jax, indexing, string-index |
| lax.while_loop: body_fun and cond_fun arguments should be… | validation | error | jax, while-loop, typeerror, callable, api-misuse |
| loading from a block pointer is not supported | exception | error | jax, triton, pallas, load, block-pointer, not-implemented |
| mxu_id must be in | validation | error | jax, pallas, tpu, accumulator, index-out-of-range |
| Only standard TMEM layout is supported, got | validation | error | gpu, mosaic, tcgen05, layout, tensor-memory |
| Out-of-bounds swap of | validation | error | jax, pallas, tpu, atomics, out-of-bounds |
| should have a different axis name from the TensorCoreMesh . | validation | error | jax, pallas, tpu, sparsecore, axis-naming |
| Subclasses should implement this method. | exception | error | jax, sharding, not-implemented, abstract-method |
| When LU decomposition matrix and b have the same number of… | validation | error | jax, lu-solve, shape-validation, broadcasting |
| Accumulator and LHS have incompatible shapes. Expected LHS… | validation | error | jax, pallas, tcgen05, mma, shape-mismatch |
| Accumulator dtype does not match value dtype | validation | error | jax, pallas, mosaic-gpu, wgmma, dtype-mismatch |
| expected 1D vector for x | validation | error | jax, numpy, polynomial, polyfit, shape-validation |
| Expected a single barrier, got a barrier reference with… | exception | error | jax, pallas, mosaic-gpu, barrier, async-store |
| Grid mapping with hijax index maps are not currently… | exception | error | jax, pallas, hijax, index-map, not-implemented |
| Input dtypes have no available implicit dtype promotion… | validation | error | jax, type-promotion, casting, dtype |
| invalid argument , expected one of | exception | error | jax, searchsorted, argument-validation |
| jnp.linalg.cond: input array must not be empty; got | exception | error | jax, linalg, condition-number, empty-array, edge-case |
| No tuned tiling found for (m, k, n) = | validation | error | jax, pallas, tpu, megablox, tiling, unsupported-shape |
| None of the leading dimensions in the transformed slice… | validation | error | jax, mosaic-gpu, tma, cluster-partitioning, shape-validation |
| NumPy arrays with zero strides are not supported as MLIR… | exception | error | numpy, strides, jax, mlir |
| tmem_addr_ref must be an i32 memref, got | exception | error | jax, mosaic, tmem, dtype, memref |
| block_kv must be a multiple of | exception | error | jax, pallas, tpu, splash-attention, segment-ids, block-size |
| {ctx.avals_out[0].dtype} | exception | error | jax, pallas, tpu, matmul, dtype, complex |
| Custom JVP rule must produce primal and tangent outputs… | exception | error | jax, custom-jvp, tangent-mismatch |
| Custom-partitioned function | exception | error | jax, gspmd, shardy, custom-partitioning, config |
| Each element of ArrayMapping must be a str or… | exception | error | jax, shardy, sharding-rule, typeerror |
| Factor names have to start with a letter, but got | exception | error | jax, shardy, sharding-rule, validation |
| Folding tiled dimensions into untiled dimensions is not… | exception | error | jax, pallas, mosaic-gpu, reshape, tiling |
| method argument to `gamma` must be one of | validation | error | jax, random, gamma, method, input-validation |
| Mismatched type shape | validation | error | jax, pallas, inline-mgpu, pytree, signature-mismatch |
| multi_dot: last dimension of each array must match first… | exception | error | jax, numpy, linalg, matmul, shape-mismatch |