jax-ml/jax · error · ValueError
ragged_all_to_all recv_sizes must be integer type.
Error message
ragged_all_to_all recv_sizes must be integer type.
What it means
recv_sizes for ragged_all_to_all must be an integer-dtype array; the abstract eval validates this and raises ValueError for float or other dtypes.
Source
Thrown at jax/_src/lax/parallel.py:1651
recv_sizes],
call_target_name=ir.StringAttr.get('ragged_all_to_all'),
backend_config=ir.DictAttr.get(ragged_all_to_all_attrs),
api_version=ir.IntegerAttr.get(ir.IntegerType.get_signless(32), 4),
).results
def _ragged_all_to_all_effectful_abstract_eval(
operand, output, input_offsets, send_sizes, output_offsets, recv_sizes,
axis_name, axis_index_groups
):
del operand, axis_index_groups
if not dtypes.issubdtype(input_offsets.dtype, np.integer):
raise ValueError("ragged_all_to_all input_offsets must be integer type.")
if not dtypes.issubdtype(send_sizes.dtype, np.integer):
raise ValueError("ragged_all_to_all send_sizes must be integer type.")
if not dtypes.issubdtype(output_offsets.dtype, np.integer):
raise ValueError("ragged_all_to_all output_offsets must be integer type.")
if not dtypes.issubdtype(recv_sizes.dtype, np.integer):
raise ValueError("ragged_all_to_all recv_sizes must be integer type.")
if len(input_offsets.shape) != 1 or input_offsets.shape[0] < 1:
raise ValueError(
"ragged_all_to_all input_offsets must be rank 1 with positive dimension"
" size, but got shape {}".format(input_offsets.shape)
)
if len(send_sizes.shape) != 1 or send_sizes.shape[0] < 1:
raise ValueError(
"ragged_all_to_all send_sizes must be rank 1 with positive dimension"
" size, but got shape {}".format(send_sizes.shape)
)
if len(output_offsets.shape) != 1 or output_offsets.shape[0] < 1:
raise ValueError(
"ragged_all_to_all output_offsets must be rank 1 with positive"
" dimension size, but got shape {}".format(output_offsets.shape)
)
if len(recv_sizes.shape) != 1 or recv_sizes.shape[0] < 1:
raise ValueError(
"ragged_all_to_all recv_sizes must be rank 1 with positive dimension"View on GitHub (pinned to 1e1c6a8fc0)
Solutions
- Cast recv_sizes to np.int64
- Ensure all four arrays (input/output offsets, send/recv sizes) are integer dtype via a pre-call check
Example fix
// before lax.ragged_all_to_all(x, out, offs, send, 'i', recv_sizes=recv.astype(jnp.float32)) // after lax.ragged_all_to_all(x, out, offs, send, 'i', recv_sizes=recv.astype(jnp.int32))
Defensive patterns
Strategy: validation
Validate before calling
assert np.issubdtype(np.asarray(recv_sizes).dtype, np.integer), 'recv_sizes must be int'
Type guard
def is_int_array(a): return np.issubdtype(np.asarray(a).dtype, np.integer)
Prevention
- Add a single validator covering all four offset/size arrays before calling ragged_all_to_all
When it happens
Trigger: Passing float-typed recv_sizes to lax.ragged_all_to_all.
Common situations: Symmetric with send/output offsets: sizes computed as floats from division or averages.
Related errors
- ragged_all_to_all input_offsets must be integer type.
- ragged_all_to_all send_sizes must be integer type.
- ragged_all_to_all output_offsets must be integer type.
- primal and tangent arguments to jax.jvp do not match; dtypes
- unexpected JAX type (e.g. shape/dtype) for gradient ref pass
AI-assisted analysis of jax-ml/jax@1e1c6a8fc0 (2026-08-27).
Data as JSON: /api/errors/284e32ace1a7bb2f.
Report an issue: GitHub.