{"record":{"id":"5950527096ef95a6","repo":"jax-ml/jax","slug":"a-scale-and-b-scale-must-both-be-present-or-absent","errorCode":null,"errorMessage":"a_scale and b_scale must both be present or absent.","messagePattern":"a_scale and b_scale must both be present or absent\\.","errorType":"validation","errorClass":"ValueError","httpStatus":null,"severity":"error","filePath":"jax/_src/pallas/mosaic_gpu/primitives.py","lineNumber":2488,"sourceCode":"        acc.transforms)\n    acc = acc.ref\n  else:\n    acc_transforms_leaves, acc_transforms_tree = [], None\n\n  if isinstance(a, pallas_core.TransformedRef):\n    a_transforms_leaves, a_transforms_tree = jax.tree.flatten(a.transforms)\n    a = a.ref\n  else:\n    a_transforms_leaves, a_transforms_tree = [], None\n\n  if isinstance(b, pallas_core.TransformedRef):\n    b_transforms_leaves, b_transforms_tree = jax.tree.flatten(b.transforms)\n    b = b.ref\n  else:\n    b_transforms_leaves, b_transforms_tree = [], None\n\n  if (is_scaled := a_scale is not None) != (b_scale is not None):\n    raise ValueError(\"a_scale and b_scale must both be present or absent.\")\n  scales = []\n  if isinstance(a_scale, pallas_core.TransformedRef):\n    a_scale_transforms_leaves, a_scale_transforms_tree = jax.tree.flatten(\n        a_scale.transforms\n    )\n    scales.append(a_scale.ref)\n  else:\n    a_scale_transforms_leaves, a_scale_transforms_tree = [], None\n    scales.append(a_scale)\n  if isinstance(b_scale, pallas_core.TransformedRef):\n    b_scale_transforms_leaves, b_scale_transforms_tree = jax.tree.flatten(\n        b_scale.transforms\n    )\n    scales.append(b_scale.ref)\n  else:\n    b_scale_transforms_leaves, b_scale_transforms_tree = [], None\n    scales.append(b_scale)\n  if not is_scaled:","sourceCodeStart":2470,"sourceCodeEnd":2506,"githubUrl":"https://github.com/jax-ml/jax/blob/1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb/jax/_src/pallas/mosaic_gpu/primitives.py#L2470-L2506","documentation":"tcgen05.mma requires a_scale and b_scale to be passed together (block-scaled MMA) or omitted together. Passing only one raises this ValueError.","triggerScenarios":"Calling tcgen05.mma(a, b, acc, k_dim=k, a_scale=s) without b_scale, or vice versa, when experimenting with MXFP8/MXFP4 block scaling.","commonSituations":"Partially wiring up scale refs in a scaled-attention kernel; refactoring scale plumbing and dropping one argument; conditional code paths that supply scales asymmetrically.","solutions":["Pass both a_scale and b_scale, or neither","If only one operand needs scaling conceptually, pass an all-ones scale ref for the other side","Gate both scale arguments on the same condition (e.g. 'if use_mxfp8: ... both')"],"exampleFix":"# before\ntcgen05.mma(a, b, acc, k_dim=k, a_scale=a_s)\n# after\ntcgen05.mma(a, b, acc, k_dim=k, a_scale=a_s, b_scale=b_s)\n# or omit both scales entirely","handlingStrategy":"validation","validationCode":"if (a_scale is None) != (b_scale is None):\n    raise ValueError('need both scales or none')\ntcgen05.mma(a, b, acc, k_dim=k, a_scale=a_scale, b_scale=b_scale)","typeGuard":null,"tryCatchPattern":null,"preventionTips":["Pass scales as a single optional tuple (a_scale, b_scale) so they cannot diverge","Gate both on one boolean"],"tags":["jax","pallas","tcgen05","block-scaling","argument-validation"],"backgroundTag":"paired-argument-validation","analyzedSha":"1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb","analyzedAt":"2026-08-27T09:53:25.647Z","schemaVersion":2},"datasetVersion":"2026-08-27T13:17:12.746Z"}