jax-ml/jax · error · TypeError
ragged_dot_general requires that group_sizes.dtype is subtyp
Error message
ragged_dot_general requires that group_sizes.dtype is subtype of np.integer.
What it means
The dtype rule for ragged_dot_general requires group_sizes to be an integer array. Float group sizes (e.g. from softmax weights or averaged counts) are rejected because group sizes drive indexing/cumsum arithmetic.
Source
Thrown at jax/_src/lax/lax.py:6467
)
if mode == RaggedDotMode.RAGGED_CONTRACTING:
out_shape = (num_groups,) + out_shape
return out_shape
def _ragged_dot_general_dtype_rule(
lhs: Array,
rhs: Array,
group_sizes: Array,
*,
ragged_dot_dimension_numbers: RaggedDotDimensionNumbers,
precision,
preferred_element_type: DTypeLike | None,
group_offset,
out_sharding,
) -> np.dtype:
if not dtypes.issubdtype(group_sizes.dtype, np.integer):
raise TypeError(
'ragged_dot_general requires that '
'group_sizes.dtype is subtype of np.integer.'
)
# defer the output dtype to dot_general, which is part of the _ragged_dot_general_impl.
return _dot_general_dtype_rule(
lhs,
rhs,
dimension_numbers=ragged_dot_dimension_numbers.dot_dimension_numbers,
precision=precision,
preferred_element_type=preferred_element_type,
out_sharding=None,
name='lax.ragged_dot_general',
)
def _ragged_dot_general_jvp_rule(
primals, tangents, ragged_dot_dimension_numbers,
precision, preferred_element_type, group_offset, out_shardingView on GitHub (pinned to 1e1c6a8fc0)
Solutions
- Cast group_sizes to an integer dtype: group_sizes.astype(jnp.int32)
- Produce counts with integer ops from the start (e.g. jnp.sum of one-hot routing, bincount)
- Validate dtypes before the call (see type guard below)
Example fix
// before group_sizes = counts.mean(axis=0) # float32 out = ragged_dot_general(x, w, group_sizes, dn, mode=mode) // after group_sizes = counts.sum(axis=0).astype(jnp.int32) out = ragged_dot_general(x, w, group_sizes, dn, mode=mode)
Defensive patterns
Strategy: type-guard
Validate before calling
group_sizes = jnp.asarray(group_sizes)
if not jnp.issubdtype(group_sizes.dtype, jnp.integer):
group_sizes = group_sizes.astype(jnp.int32) Type guard
def is_int_array(a) -> bool:
return dtypes.issubdtype(np.asarray(a).dtype, np.integer) Prevention
- Always .astype(jnp.int32) counts derived from float math
- Prefer bincount/sum for counts
When it happens
Trigger: Passing a group_sizes array with a floating dtype (e.g. result of jnp.mean, jnp.float32 counts, or a tracer carrying float dtype) to ragged_dot_general.
Common situations: Computing token-per-expert counts with float ops (softmax, division) and forgetting to cast; loading counts from a float numpy array or CSV.
Related errors
- {} does not accept dtype {}. Accepted dtypes are subtypes of
- {name} does not accept dtype {dtype_to_string(aval.dtype)}.
- {} does not accept dtype {} at position {}. Accepted dtypes
- Input type is incompatible with `preferred_element_type`. Th
- `preferred_element_type` must have the same signedness as th
AI-assisted analysis of jax-ml/jax@1e1c6a8fc0 (2026-08-27).
Data as JSON: /api/errors/26deec772604e08a.
Report an issue: GitHub.