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

  1. Pick a different index for the rhs group dim that is not in rhs_contracting or rhs_batch
  2. 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

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


AI-assisted analysis of jax-ml/jax@1e1c6a8fc0 (2026-08-27). Data as JSON: /api/errors/64061e63543f6879. Report an issue: GitHub.