{"record":{"id":"8b637a62a77a9965","repo":"jax-ml/jax","slug":"unimplemented-group-offset-support","errorCode":null,"errorMessage":"Unimplemented group_offset support.","messagePattern":"Unimplemented group_offset support\\.","errorType":"exception","errorClass":"NotImplementedError","httpStatus":null,"severity":"error","filePath":"jax/_src/lax/lax.py","lineNumber":6490,"sourceCode":"  return _dot_general_dtype_rule(\n      lhs,\n      rhs,\n      dimension_numbers=ragged_dot_dimension_numbers.dot_dimension_numbers,\n      precision=precision,\n      preferred_element_type=preferred_element_type,\n      out_sharding=None,\n      name='lax.ragged_dot_general',\n  )\n\n\ndef _ragged_dot_general_jvp_rule(\n    primals, tangents, ragged_dot_dimension_numbers,\n    precision, preferred_element_type, group_offset, out_sharding\n):\n  # note - we could ostensibly just get this by passing on the\n  # value to ragged_dot below, but, this feels cleaner.\n  if group_offset is not None:\n    raise NotImplementedError('Unimplemented group_offset support.')\n  x, y, gs = primals\n  dx, dy, _ = tangents  # no tan on the gs\n\n  # primal\n  primal_out = ragged_dot_general(\n      x,\n      y,\n      gs,\n      ragged_dot_dimension_numbers=ragged_dot_dimension_numbers,\n      precision=precision,\n      preferred_element_type=preferred_element_type,\n  )\n\n  # tangent\n  dx_out = (\n      ragged_dot_general(\n          dx,\n          y,","sourceCodeStart":6472,"sourceCodeEnd":6508,"githubUrl":"https://github.com/jax-ml/jax/blob/1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb/jax/_src/lax/lax.py#L6472-L6508","documentation":"The JVP (forward-mode autodiff) rule for ragged_dot_general does not support the group_offset argument; it only runs when group_offset is None. Passing a non-None group_offset under jvp/forward-mode raises NotImplementedError.","triggerScenarios":"Using jax.jvp (directly or via libraries that use forward-mode, e.g. jax.checkpoint in some modes or IMC/forward-over-reverse setups) on a function that calls ragged_dot_general with group_offset set.","commonSituations":"Using group_offset to resume from a previous batch segment (multi-step MoE dispatch) and then differentiating with jvp; upgrading code that worked without offsets.","solutions":["Drop the group_offset argument (pass None) when computing forward-mode derivatives","Use reverse-mode (jax.grad) instead, if the transpose rule supports group_offset in your JAX version","Compute the jvp manually by splitting the computation at offset boundaries"],"exampleFix":"// before\nout = jax.jvp(lambda x: ragged_dot_general(x, w, gs, dn, mode=mode, group_offset=off), (x,), (dx,))\n// after\nout = jax.jvp(lambda x: ragged_dot_general(x, w, gs, dn, mode=mode), (x,), (dx,))","handlingStrategy":"fallback","validationCode":"if group_offset is not None and using_jvp:\n    group_offset = None  # restructure instead","typeGuard":null,"tryCatchPattern":"except NotImplementedError as e:\n    if 'group_offset' in str(e): group_offset = None; recompute()","preventionTips":["Don't differentiate paths that pass group_offset","Reserve group_offset for inference-only code"],"tags":["jax","ragged-dot","autodiff","not-implemented"],"backgroundTag":"unsupported-autodiff-operation","analyzedSha":"1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb","analyzedAt":"2026-08-27T09:53:25.647Z","schemaVersion":2},"datasetVersion":"2026-08-27T13:17:12.746Z"}