jax-ml/jax · error · TypeError
lhs_ragged_dim {lhs_ragged_dim} not found in lhs_noncontract
Error message
lhs_ragged_dim {lhs_ragged_dim} not found in lhs_noncontracting {lhs_noncontracting}, lhs_contracting {lhs_contracting}, or lhs_batch {lhs_batch}. What it means
Raised when classifying the ragged-dot mode: lhs_ragged_dim is not a member of lhs non-contracting, contracting, or batch axes lists, so JAX cannot tell whether the ragged dimension is Mode 1/2/3.
Source
Thrown at jax/_src/lax/lax.py:6320
RAGGED_CONTRACTING = 2 # [b,m,k], [b,k,n], [b,g] -> [g,b,m,n]
RAGGED_BATCH = 3 # [b,m,k], [b,k,n], [g] -> [b,m,n]
def _ragged_dot_mode_and_dim(
lhs_rank: int, ragged_dot_dimension_numbers: RaggedDotDimensionNumbers
) -> tuple[RaggedDotMode, int]:
assert len(ragged_dot_dimension_numbers.lhs_ragged_dimensions) == 1
lhs_ragged_dim = ragged_dot_dimension_numbers.lhs_ragged_dimensions[0]
(lhs_contracting, _), (lhs_batch, _) = ragged_dot_dimension_numbers.dot_dimension_numbers
lhs_noncontracting = remaining(range(lhs_rank), lhs_contracting, lhs_batch)
if lhs_ragged_dim in lhs_noncontracting:
mode = RaggedDotMode.RAGGED_NONCONTRACTING
elif lhs_ragged_dim in lhs_contracting:
mode = RaggedDotMode.RAGGED_CONTRACTING
elif lhs_ragged_dim in lhs_batch:
mode = RaggedDotMode.RAGGED_BATCH
else:
raise TypeError(
f'lhs_ragged_dim {lhs_ragged_dim} not found in '
f'lhs_noncontracting {lhs_noncontracting}, '
f'lhs_contracting {lhs_contracting}, or '
f'lhs_batch {lhs_batch}.'
)
return mode, lhs_ragged_dim
def _ragged_dot_mode(
lhs_rank: int, ragged_dot_dimension_numbers: RaggedDotDimensionNumbers
) -> RaggedDotMode:
return _ragged_dot_mode_and_dim(lhs_rank, ragged_dot_dimension_numbers)[0]
def _is_ragged_contracting(
lhs_rank: int, ragged_dot_dimension_numbers: RaggedDotDimensionNumbers
) -> bool:
return (View on GitHub (pinned to 1e1c6a8fc0)
Solutions
- Make lhs_ragged_dim one of the lhs axes listed in batch/contracting/non-contracting sets
- Rebuild dimension numbers from scratch for the current lhs shape rather than editing an old tuple
Defensive patterns
Strategy: validation
Validate before calling
(lc, rc), (lb, rb) = rdn.dot_dimension_numbers
covered = set(lb) | set(lc)
covered |= set(range(lhs.ndim)) - set(lc) - set(lb) - {i for i in range(lhs.ndim)}
noncontracting = set(range(lhs.ndim)) - set(lc) - set(lb)
assert rdn.lhs_ragged_dimensions[0] in noncontracting | set(lc) | set(lb) Prevention
- Regenerate RaggedDotDimensionNumbers whenever lhs rank changes
- Keep ragged dim derived from a named constant, not a magic index
When it happens
Trigger: Calling jax.lax.ragged_dot_general with ragged_dot_dimension_numbers whose lhs_ragged_dimensions[0] is an index not covered by the dot dimension numbers (e.g. >= lhs.ndim or only present in rhs specs).
Common situations: Hand-building RaggedDotDimensionNumbers; desync between dot_dimension_numbers and the ragged dim after refactoring shapes or axes.
Related errors
- ragged_dot_general expects exactly one lhs ragged dimension.
- No 4+ dimensional dimension_number defaults.
- convolution dimension_numbers list/tuple must be length 3, g
- convolution dimension_numbers elements must be strings, got
- convolution dimension_numbers[{}] must have len equal to the
AI-assisted analysis of jax-ml/jax@1e1c6a8fc0 (2026-08-27).
Data as JSON: /api/errors/359be4a31e550de5.
Report an issue: GitHub.