{"record":{"id":"359be4a31e550de5","repo":"jax-ml/jax","slug":"lhs-ragged-dim-lhs-ragged-dim-not-found-in-lhs-n","errorCode":null,"errorMessage":"lhs_ragged_dim {lhs_ragged_dim} not found in lhs_noncontracting {lhs_noncontracting}, lhs_contracting {lhs_contracting}, or lhs_batch {lhs_batch}.","messagePattern":"lhs_ragged_dim (.+?) not found in lhs_noncontracting (.+?), lhs_contracting (.+?), or lhs_batch (.+?)\\.","errorType":"validation","errorClass":"TypeError","httpStatus":null,"severity":"error","filePath":"jax/_src/lax/lax.py","lineNumber":6320,"sourceCode":"  RAGGED_CONTRACTING = 2  #    [b,m,k], [b,k,n],   [b,g] -> [g,b,m,n]\n  RAGGED_BATCH = 3  #          [b,m,k], [b,k,n],   [g]   -> [b,m,n]\n\n\ndef _ragged_dot_mode_and_dim(\n    lhs_rank: int, ragged_dot_dimension_numbers: RaggedDotDimensionNumbers\n) -> tuple[RaggedDotMode, int]:\n  assert len(ragged_dot_dimension_numbers.lhs_ragged_dimensions) == 1\n  lhs_ragged_dim = ragged_dot_dimension_numbers.lhs_ragged_dimensions[0]\n  (lhs_contracting, _), (lhs_batch, _) = ragged_dot_dimension_numbers.dot_dimension_numbers\n  lhs_noncontracting = remaining(range(lhs_rank), lhs_contracting, lhs_batch)\n  if lhs_ragged_dim in lhs_noncontracting:\n    mode = RaggedDotMode.RAGGED_NONCONTRACTING\n  elif lhs_ragged_dim in lhs_contracting:\n    mode = RaggedDotMode.RAGGED_CONTRACTING\n  elif lhs_ragged_dim in lhs_batch:\n    mode = RaggedDotMode.RAGGED_BATCH\n  else:\n    raise TypeError(\n        f'lhs_ragged_dim {lhs_ragged_dim} not found in '\n        f'lhs_noncontracting {lhs_noncontracting}, '\n        f'lhs_contracting {lhs_contracting}, or '\n        f'lhs_batch {lhs_batch}.'\n    )\n  return mode, lhs_ragged_dim\n\n\ndef _ragged_dot_mode(\n    lhs_rank: int, ragged_dot_dimension_numbers: RaggedDotDimensionNumbers\n) -> RaggedDotMode:\n  return _ragged_dot_mode_and_dim(lhs_rank, ragged_dot_dimension_numbers)[0]\n\n\ndef _is_ragged_contracting(\n    lhs_rank: int, ragged_dot_dimension_numbers: RaggedDotDimensionNumbers\n) -> bool:\n  return (","sourceCodeStart":6302,"sourceCodeEnd":6338,"githubUrl":"https://github.com/jax-ml/jax/blob/1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb/jax/_src/lax/lax.py#L6302-L6338","documentation":"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.","triggerScenarios":"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).","commonSituations":"Hand-building RaggedDotDimensionNumbers; desync between dot_dimension_numbers and the ragged dim after refactoring shapes or axes.","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"],"exampleFix":null,"handlingStrategy":"validation","validationCode":"(lc, rc), (lb, rb) = rdn.dot_dimension_numbers\ncovered = set(lb) | set(lc)\ncovered |= set(range(lhs.ndim)) - set(lc) - set(lb) - {i for i in range(lhs.ndim)}\nnoncontracting = set(range(lhs.ndim)) - set(lc) - set(lb)\nassert rdn.lhs_ragged_dimensions[0] in noncontracting | set(lc) | set(lb)","typeGuard":null,"tryCatchPattern":null,"preventionTips":["Regenerate RaggedDotDimensionNumbers whenever lhs rank changes","Keep ragged dim derived from a named constant, not a magic index"],"tags":["jax","ragged-dot-general","dimension-numbers"],"backgroundTag":"invalid-dimension-numbers","analyzedSha":"1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb","analyzedAt":"2026-08-27T09:53:25.647Z","schemaVersion":2},"datasetVersion":"2026-08-27T13:17:12.746Z"}