{"record":{"id":"d728e3b3c82fdf86","repo":"jax-ml/jax","slug":"unsupported-transforms-for-ref-transforms-tran","errorCode":null,"errorMessage":"Unsupported transforms for {ref}. Transforms {transforms}.","messagePattern":"Unsupported transforms for (.+?)\\. Transforms (.+?)\\.","errorType":"exception","errorClass":"NotImplementedError","httpStatus":null,"severity":"error","filePath":"jax/_src/pallas/mosaic_gpu/primitives.py","lineNumber":2969,"sourceCode":"      acc_transforms_leaves_avals,\n      a_transforms_leaves_avals,\n      b_transforms_leaves_avals,\n      _,\n      a_scale_transforms_leaves_avals,\n      b_scale_transforms_leaves_avals,\n      a_sparse_metadata_transforms_leaves_avals,\n  ) = transforms_avals_lists\n\n  def handle_transforms_and_get_ref(tree, leaves, leaves_avals, ref, ref_aval, handle_transposes=True):\n    if tree is None:\n      return ref\n    transforms = tree.unflatten(leaves)\n    transform_avals = tree.unflatten(leaves_avals)\n    ref, _, transforms = lowering._handle_transforms(\n        ctx, ref_aval, ref, transform_avals, transforms, handle_transposes=handle_transposes\n    )\n    if transforms:\n      raise NotImplementedError(\n          f\"Unsupported transforms for {ref}. Transforms {transforms}.\"\n      )\n    return ref\n\n  acc_ref = handle_transforms_and_get_ref(\n      acc_transforms_tree,\n      acc_transforms_leaves,\n      acc_transforms_leaves_avals,\n      acc_ref,\n      acc_aval,\n      handle_transposes=False,\n  )\n\n  a_ref = handle_transforms_and_get_ref(\n      a_transforms_tree,\n      a_transforms_leaves,\n      a_transforms_leaves_avals,\n      a_ref,","sourceCodeStart":2951,"sourceCodeEnd":2987,"githubUrl":"https://github.com/jax-ml/jax/blob/1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb/jax/_src/pallas/mosaic_gpu/primitives.py#L2951-L2987","documentation":"The warp-group variant of the tcgen05 MMA lowering routes the accumulator and operand references through _handle_transforms; if any transforms remain unhandled (notably transposes when handle_transposes cannot apply), it raises NotImplementedError naming the offending ref and transforms.","triggerScenarios":"Using the warp-group tcgen05 MMA lowering with a reference whose transform tree leaves residual transforms (e.g. transpose on a TMEM accumulator ref).","commonSituations":"Writing warp-group MMA kernels with transposed accumulators or exotic layouts on refs fed to the MMA.","solutions":["Remove the remaining transforms on the named ref (message shows exactly which ref and transforms)","Apply the transpose/layout change manually to the buffer contents instead of via a ref transform"],"exampleFix":"// before\nacc_ref_t = plgpu.transpose_ref(acc_ref)\ntcgen05_mma_wg(a, b, acc_ref_t)\n// after\ntcgen05_mma_wg(a, b, acc_ref)","handlingStrategy":"validation","validationCode":"def check_ref_clean(ref, transforms):\n    if transforms:\n        raise NotImplementedError(f'clean {transforms} on {ref} first')","typeGuard":null,"tryCatchPattern":null,"preventionTips":["Audit transform trees on all MMA refs before calling warp-group MMA","Use handle_transposes-compatible layouts"],"tags":["jax","pallas","tcgen05","warp-group","transforms","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"}