{"record":{"id":"f87ecd2e1cc3670f","repo":"jax-ml/jax","slug":"ragged-dot-vmap-over-any-dim-but-0-nyi","errorCode":null,"errorMessage":"ragged_dot vmap over any dim but 0 - NYI","messagePattern":"ragged_dot vmap over any dim but 0 - NYI","errorType":"exception","errorClass":"NotImplementedError","httpStatus":null,"severity":"error","filePath":"jax/_src/lax/lax.py","lineNumber":6633,"sourceCode":"      if ad.is_undefined_primal(y)\n      else _ragged_dot_grad(ct, y, grad_x_dims, x.aval)\n  )\n  y_bar = (\n      None\n      if ad.is_undefined_primal(x)\n      else _ragged_dot_grad(x, ct, grad_y_dims, y.aval)\n  )\n  return x_bar, y_bar, None\n\n\ndef _ragged_dot_batch_unpack_args(batched_args):\n  lhs, rhs, _ = batched_args\n  return (lhs, rhs)\n\n\ndef _ragged_dot_batch_unpack_dims(batch_dims):\n  if not all(dim == 0 for dim in batch_dims):\n    raise NotImplementedError('ragged_dot vmap over any dim but 0 - NYI')\n  lbd, rbd, _ = batch_dims\n  return (lbd, rbd)\n\n\ndef _ragged_dot_general_invoke_prim(\n    group_sizes,\n    lhs,\n    rhs,\n    new_ragged_dot_dimension_numbers,\n    precision,\n    preferred_element_type,\n    out_sharding,\n):\n  del out_sharding\n  return ragged_dot_general(\n      lhs,\n      rhs,\n      group_sizes,","sourceCodeStart":6615,"sourceCodeEnd":6651,"githubUrl":"https://github.com/jax-ml/jax/blob/1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb/jax/_src/lax/lax.py#L6615-L6651","documentation":"The vmap batch rule for ragged_dot only supports batching over dimension 0 of every batched operand. Batching over any other dimension is not yet implemented (NYI).","triggerScenarios":"Applying jax.vmap to a function using ragged_dot/ragged_dot_general where the batched axis maps to a non-zero dimension of lhs, rhs, or group_sizes.","commonSituations":"Vectorizing a per-example ragged matmul but in_axes points at dim 1 (e.g. in_axes=1 on a (seq, batch, d) tensor), or transposing inputs so the batch dim isn't leading.","solutions":["Transpose/squeeze so the batched dimension is axis 0 before vmap (e.g. x.transpose(1,0,2) or in_axes=0)","Use in_axes=0 (or None) for all ragged_dot operands","Fall back to a manual loop or lax.map over the batch"],"exampleFix":"// before\nf = jax.vmap(ragged_step, in_axes=(1, None, None))  # x batched on dim 1\n// after\nx2 = x.transpose(1, 0, 2)\nf = jax.vmap(ragged_step, in_axes=(0, None, None))","handlingStrategy":"validation","validationCode":"assert all(ax in (0, None) for ax in in_axes), 'batch on dim 0 only'","typeGuard":null,"tryCatchPattern":null,"preventionTips":["Keep the batch axis leading in all tensors feeding ragged dot","Transposing inside vmapped fn, not via in_axes"],"tags":["jax","ragged-dot","vmap","batching","not-implemented"],"backgroundTag":"vmap-unsupported-batching","analyzedSha":"1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb","analyzedAt":"2026-08-27T09:53:25.647Z","schemaVersion":2},"datasetVersion":"2026-08-27T13:17:12.746Z"}