{"record":{"id":"d6653170a6337cac","repo":"jax-ml/jax","slug":"unimplemented-transforms-for-tmem-refs-transform","errorCode":null,"errorMessage":"Unimplemented transforms for TMEM refs. {transforms=}","messagePattern":"Unimplemented transforms for TMEM refs\\. (.+?)","errorType":"validation","errorClass":"NotImplementedError","httpStatus":null,"severity":"error","filePath":"jax/_src/pallas/mosaic_gpu/primitives.py","lineNumber":4107,"sourceCode":"def _async_load_tmem_lowering_rule(\n    ctx: lowering.LoweringRuleContext,\n    x_ref,\n    *leaves,\n    tree,\n    reduce: Literal[\"max\", \"min\", \"absmax\", \"absmin\"] | None = None,\n):\n  assert isinstance(x_ref, tcgen05.TMEMRef)\n  x_aval = ctx.avals_in[0]\n  assert isinstance(x_aval, state_types.AbstractRef)\n  transforms = jax.tree.unflatten(tree, leaves)\n  transform_avals = tree.unflatten(\n      ctx.avals_in[1 : 1 + tree.num_leaves]\n  )\n  x_tmem, _, transforms = lowering._handle_transforms(\n      ctx, x_aval, x_ref, transform_avals, transforms, handle_transposes=False,\n      handle_reshapes=False)\n  if transforms:\n    raise NotImplementedError(\n        f\"Unimplemented transforms for TMEM refs. {transforms=}\"\n    )\n  layout_hint = None\n  if isinstance(ctx.out_layout_hint, mgpu.TiledLayout):\n    layout_hint = ctx.out_layout_hint\n  is_signed = mgpu_utils.is_signed(ctx.avals_out[0].dtype)\n  res = x_tmem.load(layout=layout_hint, is_signed=is_signed, reduce=reduce)\n  return (res,) if reduce is None else res\n\n\n@lowering.register_lowering_rule(\n    async_load_tmem_p, mgpu.LoweringSemantics.Warpgroup\n)\ndef _async_load_tmem_lowering_rule_wg(\n    ctx: lowering.LoweringRuleContext,\n    x_ref: ir.Value,\n    *leaves,\n    tree,","sourceCodeStart":4089,"sourceCodeEnd":4125,"githubUrl":"https://github.com/jax-ml/jax/blob/1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb/jax/_src/pallas/mosaic_gpu/primitives.py#L4089-L4125","documentation":"When lowering a TMEM operation, residual transforms remained after _handle_transforms with transposes and reshapes disabled; TMEM refs do not support those transforms. Any leftover transform raises NotImplementedError.","triggerScenarios":"Applying a transpose or reshape (or any transform) to a TMEM ref before an operation lowered with handle_transposes=False and handle_reshapes=False.","commonSituations":"Writing tmem_ref.T or reshaping a TMEM ref; assuming TMEM behaves like SMEM regarding transforms.","solutions":["Remove the transpose/reshape on the TMEM ref; do the transpose after loading into registers","Allocate the TMEM ref with the final desired shape/layout instead of transforming it"],"exampleFix":"// before\nx = load(tmem_ref.T)\n// after\nx = load(tmem_ref).T","handlingStrategy":"validation","validationCode":"assert not transforms, f'leftover TMEM transforms: {transforms}'","typeGuard":null,"tryCatchPattern":null,"preventionTips":["Never transpose or reshape TMEM refs; transform values after load"],"tags":["jax","pallas","tmem","transforms","not-implemented"],"backgroundTag":"unsupported-transform-combination","analyzedSha":"1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb","analyzedAt":"2026-08-27T09:53:25.647Z","schemaVersion":2},"datasetVersion":"2026-08-27T13:17:12.746Z"}