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

  1. Make lhs_ragged_dim one of the lhs axes listed in batch/contracting/non-contracting sets
  2. 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

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


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