{"record":{"id":"1eeda67d60ef78ae","repo":"jax-ml/jax","slug":"unsupported-transforms-a-scale-transforms","errorCode":null,"errorMessage":"Unsupported transforms: {a_scale_transforms}","messagePattern":"Unsupported transforms: (.+?)","errorType":"exception","errorClass":"NotImplementedError","httpStatus":null,"severity":"error","filePath":"jax/_src/pallas/mosaic_gpu/primitives.py","lineNumber":2804,"sourceCode":"    accumulate = mgpu.c(accumulate, ir.IntegerType.get_signless(1))\n  elif isinstance(accumulate, mgpu.FragmentedArray):\n    accumulate = accumulate.registers.item()\n    assert isinstance(accumulate, ir.Value)\n\n  if a_scale_ref is not None and a_scale_transforms_tree is not None:\n    assert isinstance(a_scale_ref_aval, state.AbstractRef)\n    a_scale_transforms = a_scale_transforms_tree.unflatten(\n        a_scale_transforms_leaves\n    )\n    a_scale_transform_avals = a_scale_transforms_tree.unflatten(\n        a_scale_transforms_leaves_avals\n    )\n    a_scale_ref, _, a_scale_transforms = lowering._handle_transforms(\n        ctx, a_scale_ref_aval, a_scale_ref, a_scale_transform_avals,\n        a_scale_transforms\n    )\n    if a_scale_transforms:\n      raise NotImplementedError(\n          f\"Unsupported transforms: {a_scale_transforms}\"\n      )\n  if b_scale_ref is not None and b_scale_transforms_tree is not None:\n    assert isinstance(b_scale_ref_aval, state.AbstractRef)\n    b_scale_transforms = b_scale_transforms_tree.unflatten(\n        b_scale_transforms_leaves\n    )\n    b_scale_transform_avals = b_scale_transforms_tree.unflatten(\n        b_scale_transforms_leaves_avals\n    )\n    b_scale_ref, _, b_scale_transforms = lowering._handle_transforms(\n        ctx, b_scale_ref_aval, b_scale_ref, b_scale_transform_avals,\n        b_scale_transforms\n    )\n    if b_scale_transforms:\n      raise NotImplementedError(f\"Unsupported transforms: {b_scale_transforms}\")\n  if a_sparse_metadata_transforms_tree is not None:\n    a_sparse_metadata_transforms = a_sparse_metadata_transforms_tree.unflatten(","sourceCodeStart":2786,"sourceCodeEnd":2822,"githubUrl":"https://github.com/jax-ml/jax/blob/1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb/jax/_src/pallas/mosaic_gpu/primitives.py#L2786-L2822","documentation":"When using scaled MMA (fp8/fp4 block scaling on the A operand), any layout transforms (e.g. transposes) left over on the a_scale reference after _handle_transforms cannot be lowered, so the lowering raises NotImplementedError.","triggerScenarios":"Passing a_scale_ref to tcgen05_mma together with transform trees (e.g. a transpose transform) that apply to the A-scale reference and cannot be handled during lowering.","commonSituations":"Building fp8 block-scaled GEMM kernels where the scale tensors are loaded from transposed or transformed references.","solutions":["Remove the transforms on the a_scale reference (load it un-transformed; materialize the transpose manually)","Pre-transpose/rearrange the scale tensor in SMEM before passing it to the MMA"],"exampleFix":"// before\ntcgen05_mma(a, b, acc, a_scale=scale_ref, a_scale_transforms=transforms_with_transpose)\n// after\ntcgen05_mma(a, b, acc, a_scale=scale_ref)  # scale loaded already in the right layout","handlingStrategy":"validation","validationCode":"assert not a_scale_transforms or all(t is None for t in a_scale_transforms)","typeGuard":null,"tryCatchPattern":null,"preventionTips":["Load scale tensors with the exact final layout","Avoid applying BlockSpec transforms to scale refs"],"tags":["jax","pallas","tcgen05","fp8","scaling","not-implemented"],"backgroundTag":"unsupported-transform-on-operand","analyzedSha":"1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb","analyzedAt":"2026-08-27T09:53:25.647Z","schemaVersion":2},"datasetVersion":"2026-08-27T13:17:12.746Z"}