jax-ml/jax · error · TypeError
ragged_dot_general requires rhs group dimension numbers to b
Error message
ragged_dot_general requires rhs group dimension numbers to be distinct from contracting and batch dimensions.
What it means
The rhs group dimension must be a distinct dimension not also used as a contracting or batch dimension. Reusing the same index in multiple roles makes the dot semantics ambiguous, so JAX rejects it.
Source
Thrown at jax/_src/lax/lax.py:6432
# Validate properties of the rhs group dimension(s).
rhs_group_dims = ragged_dot_dimension_numbers.rhs_group_dimensions
match mode:
case RaggedDotMode.RAGGED_CONTRACTING | RaggedDotMode.RAGGED_BATCH:
if len(rhs_group_dims) != 0:
raise TypeError(
'ragged_dot_general requires zero group dimensions in the rhs '
'when lhs ragged dimension is contracting or batch.'
)
case RaggedDotMode.RAGGED_NONCONTRACTING:
if len(rhs_group_dims) != 1:
raise TypeError(
'ragged_dot_general requires exactly one rhs group dimension '
'when lhs ragged dimension is noncontracting.'
)
rhs_group_dim = rhs_group_dims[0]
_check_in_range(rhs_group_dim, rhs.ndim, 'rhs group dimension', 'rhs')
if rhs_group_dim in rhs_batch or rhs_group_dim in rhs_contracting:
raise TypeError(
'ragged_dot_general requires rhs group dimension numbers to be '
'distinct from contracting and batch dimensions.'
)
if rhs.shape[rhs_group_dim] != num_groups:
raise TypeError(
'expected rhs group dimension size to be '
f'{num_groups}, got {rhs.shape[rhs_group_dim]}.'
)
out_shape = _dot_general_shape_rule(
lhs,
rhs,
dimension_numbers=ragged_dot_dimension_numbers,
precision=precision,
preferred_element_type=preferred_element_type,
out_sharding=None,
)
if mode == RaggedDotMode.RAGGED_CONTRACTING:View on GitHub (pinned to 1e1c6a8fc0)
Solutions
- Pick a different index for the rhs group dim that is not in rhs_contracting or rhs_batch
- Print/assert your dimension numbers before the call: assert len({*rhs_group, *rhs_contract, *rhs_batch}) == len(rhs_group)+len(rhs_contract)+len(rhs_batch)
Example fix
// before dn = RaggedDotDimensionNumbers((2,),(1,),(0,),(0,),(1,)) // group=1 == contract=1 // after dn = RaggedDotDimensionNumbers((2,),(1,),(0,),(0,),(0,)) // group=0, distinct
Defensive patterns
Strategy: validation
Validate before calling
used = set(dn.rhs_contracting) | set(dn.rhs_batch) assert all(g not in used for g in dn.rhs_group_dimensions)
Prevention
- Uniqueness-check all rhs index lists before the call
- Log dimension numbers once at model-build time
When it happens
Trigger: Passing RaggedDotDimensionNumbers in RAGGED_NONCONTRACTING mode where rhs_group_dimensions[0] also appears in rhs contracting dimensions or rhs batch dimensions.
Common situations: Hand-building dimension-number tuples and accidentally duplicating an index (e.g. contract=(1,), group=(1,)) when adapting a plain dot_general dimension_numbers to the ragged variant.
Related errors
- convolution dimension_numbers[{}] must have len equal to the
- ragged_dot_general requires zero group dimensions in the rhs
- ragged_dot_general requires exactly one rhs group dimension
- scan got `length` argument of {} which disagrees with leadin
- No 4+ dimensional dimension_number defaults.
AI-assisted analysis of jax-ml/jax@1e1c6a8fc0 (2026-08-27).
Data as JSON: /api/errors/64061e63543f6879.
Report an issue: GitHub.