jax-ml/jax
Documented errors, page 4 of 23. Back to jax-ml/jax
| Code / Message | Type | Severity | Tags |
|---|---|---|---|
| cannot cast tp | validation | error | jax, pallas, triton, dtype-cast, float8, notimplementederror |
| cond_fun must return a boolean scalar, but got pytree | validation | error | jax, while-loop, cond-fun, pytree, boolean-scalar |
| index arguments to dynamic_update_slice must be integers of… | exception | error | jax, lax, dtype, indices |
| Memory space is not supported by mesh | validation | error | jax, pallas, memory-space, mesh, gpu, tpu |
| out_specs_fn already specified | validation | error | jax, colocated-python, api-misuse, configuration |
| shard_map _specs argument must refer to an axis marked as… | exception | error | shard-map, manual-axes, partition-spec, jax |
| After moving axes to end, leading shape of a must match… | exception | error | jax, numpy, linalg, tensor, shape-validation |
| Arguments to jax.numpy.gcd must be integers. | exception | error | jax, numpy, dtype-validation, integer-required |
| Cannot use `partial_eval_jaxpr_custom` with stateful jaxprs. | validation | error | jax, remat, partial-eval, stateful-jaxpr, control-flow |
| dtype argument to `binomial` must be a float dtype, got | validation | error | jax, random, dtype-validation |
| Input arrays must have prod(a.shape[:b.ndim]) ==… | exception | error | jax, numpy, linalg, tensor, shape-validation |
| Non-trivial windowing is not supported for grid-free… | exception | error | jax, pallas, tpu, grid, windowing |
| Out-of-bounds write of | validation | error | jax, pallas, tpu, out-of-bounds, write |
| Unsupported aval type | validation | error | jax, pallas, aval, not-implemented, mlir |
| Custom JVP rule for function must produce a pair (list or… | exception | error | jax, custom-jvp, pytree-mismatch |
| input_memory_space_colors only supports HBM, VMEM and SMEM | exception | error | jax, tpu, custom-call, serialization, not-implemented |
| The input arguments to the custom_jvp-decorated function | exception | error | jax, custom-jvp, kwargs, signature-binding |
| Array slice indices must have static start/stop/step to be… | validation | error | jax, jit, tracing, dynamic-slice |
| Async copies with require the last dimension of the slice… | validation | error | jax, mosaic-gpu, tma, swizzling, shared-memory |
| Batching over dynamic grid values is not supported yet. | exception | error | jax, pallas, mosaic, vmap, batching, tpu, not-implemented |
| BSCR.from_bcoo requires n_sparse=2; got | exception | error | jax, sparse, bcsr, bcoo, format-conversion |
| convolution dimension_numbers | validation | error | jax, convolution, dimension-numbers, validation |
| Expected an `NDIndexer`, but got | exception | error | jax, mosaic-gpu, barrier, type-validation |
| lax.associative_scan: fn argument should be callable. | validation | error | jax, associative-scan, type-error, callable |
| Mismatched type | validation | error | jax, pallas, inline-mgpu, type-mismatch |
| Non-trivial indexing on WGMMAAbstractAccumulatorRef is not… | validation | error | jax, pallas, mosaic-gpu, wgmma, not-implemented, indexing |
| the bwd rule attached to produced an output of type which… | exception | error | jax, custom-vjp, autodiff, type-mismatch |
| Sharding rule has operands, but the operation has operands | validation | error | jax, sharding, custom-partitioning, mlir, sdy |
| Sum of sizes must be equal to dimension of the operand… | exception | error | jax, pallas, lax-split, shape-mismatch, validation |
| the `static_argnums` argument to `jax.checkpoint` /… | validation | error | pytree, tree-map, structure-mismatch, jax |
| unbound axis name | validation | error | jax, mesh, collectives, axis-name, sharding |
| under vmap, the of produced an output batched along the… | validation | error | jax, vmap, custom-vjp, custom-jvp, batching |
| iteration over a 0-d key array | exception | error | jax, prng, iteration, unpacking |
| num_segments must be non-negative. | validation | error | jax, segment-reduction, argument-validation, valueerror |
| out_dtype argument in binary_op_lowering_rule_wg | exception | error | jax, pallas, mosaic-gpu, dtype, lowering, not-implemented |
| _rbg_random_bits got invalid prng key. | exception | error | jax, prng, rbg, shape-mismatch |
| shardings should container 4 inputs, but got | validation | error | jax, fp8, sharding, spmd, matmul |
| Transforms not supported for matmul_acc_lhs. | exception | error | jax, pallas, tpu, matmul, not-implemented |
| bcoo_slice: input should be BCOO array, got type(mat)= | exception | error | jax, sparse, bcoo, type-error, slice |
| mapped axes must have same shape; got | exception | error | jax, scipy, signal, vmap, shape-mismatch |
| Sparse meta layout loads unsupported. | exception | error | mosaic, gpu, tcgen05, sparse, tensor-memory, not-implemented, jax |
| The input part of spec in out_sharding should match the… | validation | error | jax, nn, one-hot, sharding, named-sharding |
| Transpose of Einsum with multiple linear inputs is not… | exception | error | jax, einsum, autodiff, not-implemented |
| Array shape mismatch: expected | validation | error | jax, pallas, inline-mgpu, shape-mismatch |
| Body jaxpr has consts. If you see this error, please open… | error_code | error | jax, state-discharge, internal-error, while-loop |
| Cannot write to the same ref in both cond and body of while… | error_code | error | jax, state-api, refs, while-loop, state-discharge |
| Error reading persistent compilation cache entry for | console | warning | jax, compilation-cache, persistence, io |
| Expected start indices, got | validation | error | |
| function traced for returned a mutable array reference of… | validation | error | jax, mutable-arrays, refs, tracing, return-value |
| group_offset must be a ()-shaped array. Got | validation | error | jax, pallas, tpu, megablox, shape-validation |
| ind must be a positive integer; got | exception | error | jax, scipy, linalg, tensor, argument-validation |
| Only leading gather dimensions allowed. | validation | error | jax, mosaic-gpu, tma, gather, async-load, layout-inference |
| The input `a` must be at least a 2-D array. | exception | error | jax, polar-decomposition, input-validation, shape-error |
| unsupported dtypes: and | validation | error | jax, pallas, triton, jnp-minimum, dtype, notimplementederror |
| `weights` input should be one-dimensional. | exception | error | jax, scipy, kde, weights, shape-validation |
| AbstractMesh does not implement | validation | error | jax, mesh, abstract-class, not-implemented |
| Dict key mismatch; expected keys | validation | error | jax, pytree, dict, key-mismatch |
| Expected counter to be a scalar integer ref; got | exception | error | jax, prng, stateful-rng, ref, experimental |
| Only gathers along the two minormost dimensions supported… | exception | error | jax, pallas, tpu, tensorcore, gather |
| performing a set/swap operation with a differentiated value… | exception | error | jax, autodiff, state-primitives, jvp |
| Require q_seqlen and kv_seqlen to use packed layout | validation | error | jax, attention, varlen, packed-layout, cudnn |
| Unsupported index: of type | exception | error | pallas, mosaic-gpu, indexing, type-error |
| Can not bitcast memory region of size | exception | error | jax, pallas, mosaic-gpu, bitcast, size-mismatch, alignment |
| JAX array with PRNGKey dtype cannot be converted to a NumPy… | exception | error | jax, prng, numpy, serialization |
| Out-of-bounds masked swap of | validation | error | jax, pallas, tpu, atomics, mask, out-of-bounds |
| Value of type is not convertible to hex. | exception | error | jax, tracer, hex, formatting, debugging |
| Block shape for (= ) must have the same number of… | validation | error | jax, pallas, block-spec, shape-mismatch |
| custom_jvp-decorated function | exception | error | jax, custom-jvp, tracer-error, closure |
| get not supported yet | exception | error | jax, pallas, get, indexer, not-implemented |
| Invalid shape for `addupdate`. Ref shape | validation | error | jax, shape-mismatch, state-primitives, in-place-update |
| to_dlpack can only pack a dlpack tensor from an array on a… | exception | error | jax, dlpack, interop, sharding, multi-device |
| Unknown GPU platform for __dlpack__ | exception | error | jax, dlpack, gpu, platform-detection, version-mismatch |
| Unsupported dot algorithm | exception | error | jax, pallas, triton, precision, dot-algorithm |
| Unsupported scalar attribute type | exception | error | jax, dtype, mlir, numpy |
| Value of type is not convertible to oct. | exception | error | jax, tracer, oct, formatting |
| Accumulator and RHS have incompatible shapes. Expected RHS… | validation | error | jax, pallas, tcgen05, collective-mma, shape-mismatch |
| Cannot lower jaxpr with effects | exception | error | jax, effects, custom-partitioning, lowering, pmap |
| HLO comparison for extended dtype | exception | error | jax, extended-dtype, comparison, lax, hlo |
| Invalid mixing of symbolic scopes | exception | error | jax, shape-polymorphism, scope, export |
| must be less or equal to at sequence . | validation | error | jax, pallas, tpu, ragged-attention, length-validation |
| The kernel function in the pallas_call | exception | error | jax, pallas, kernel, return-value, validation |
| The number of sources must match the packing factor | exception | error | jax, pallas, tpu, dtype, shape-mismatch |
| the VJP function was applied before restoring its… | validation | error | jax, vjp, memory, checkpointing, saveable-args |
| unreduced cannot contain None. All elements in unreduced… | validation | error | jax, sharding, partition-spec, mesh |
| Cannot divide by . | exception | error | jax, shape-polymorphism, floor-division, export |
| Cond jaxpr has consts. If you see this error, please open… | error_code | error | jax, state-discharge, internal-error, while-loop |
| In arange with non-constant arguments all of start, stop… | validation | error | jax, arange, symbolic-dim, shape-polymorphism, valueerror |
| keyword arguments could not be resolved to positions | exception | error | jax, custom-partitioning, kwargs, typeerror |
| Unhandled transforms for multimem_store | exception | error | jax, pallas, transforms, multimem, lowering |
| BlockMapping for has captured constants | validation | error | jax, pallas, block-mapping, internal, closure |
| JAX Arrays do not implement the arr.flat property: consider… | exception | error | jax, flat, not-implemented, numpy-compat |
| Loads and stores are only allowed on VMEM and SMEM… | exception | error | jax, pallas, tpu, memory-space, async-copy |
| out_specs passed to shard_map should be equal to the… | exception | error | shard-map, reduced, out-specs, sharding-mismatch, jax |
| The return value of the policies should be a boolean. Got | validation | warning | pytree, deprecation, sets, generators, jax |
| Error refining shapes. | exception | error | jax, polymorphic-shapes, dynamic-shapes, mlir, lowering |
| For a cross-host reshard in multi-controller JAX, input and… | validation | error | jax, multi-controller, resharding, device-mismatch |
| must be in range | validation | error | jax, pallas, tpu, block-size, paged-attention |
| Only 32-, 64- and 128-bit stores are supported | validation | error | gpu, mosaic, dsmem, multimem, store, bit-width |
| Transform mismatch: got | validation | error | jax, pallas, transform-mismatch, validation |
| Context mesh cannot be empty. Please use `jax.set_mesh` API… | validation | error | jax, mesh, sharding, context |