{"record":{"id":"da3fc3b584b8d0f5","repo":"jax-ml/jax","slug":"unsupported-tmem-ref-ref","errorCode":null,"errorMessage":"Unsupported TMEM ref {ref}.","messagePattern":"Unsupported TMEM ref (.+?)\\.","errorType":"exception","errorClass":"NotImplementedError","httpStatus":null,"severity":"error","filePath":"jax/_src/pallas/mosaic_gpu/lowering.py","lineNumber":1648,"sourceCode":"          else:\n            ref_bytes = ref_bits // 8\n            ref = mgpu.memref_slice(ref, slice(offset, offset + ref_bytes))\n            ref = _handle_dtype_bitcast(\n                ref,\n                ir.MemRefType(ref.type).element_type,\n                mlir_dtype,\n            )\n            ref = mgpu.memref_reshape(ref, transformed_shape)\n        elif input_ref_ty.memory_space == mgpu_utils.tmem():\n\n          if isinstance(ref.owner, mgpu.dialect.SliceTmemOp):\n            source_slice_op = ref.owner\n          elif isinstance(\n              ref.owner, mgpu.dialect.TmemLayoutCastOp\n          ) and isinstance(ref.owner.operands[0].owner, mgpu.dialect.SliceTmemOp):\n            source_slice_op = ref.owner.operands[0].owner\n          else:\n            raise NotImplementedError(f\"Unsupported TMEM ref {ref}.\")\n\n          base_offset = source_slice_op.offset.value\n          assert isinstance(base_offset, int)  # make pyrefly happy\n          total_offset = base_offset + offset\n          ref_ty = ir.MemRefType.get(\n              transformed_shape, mlir_dtype, memory_space=mgpu_utils.tmem()\n          )\n          alloc_id = source_slice_op.alias_id\n          assert alloc_id is not None\n          # TODO(bchetioui): Use a scheme resilient to hash collisions.\n          alias_id = hash((offset, alloc_id.value, alias_group_idx))\n          slice_op = mgpu.dialect.SliceTmemOp(\n              ref_ty, source_slice_op.source, total_offset\n          )\n          i64 = ir.IntegerType.get_signless(64)\n          slice_op.attributes[\"alias_id\"] = ir.IntegerAttr.get(i64, alias_id)\n          ref = slice_op.result\n          assert layout is not None","sourceCodeStart":1630,"sourceCodeEnd":1666,"githubUrl":"https://github.com/jax-ml/jax/blob/1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb/jax/_src/pallas/mosaic_gpu/lowering.py#L1630-L1666","documentation":"For TMEM aliases, the lowering can only recover the base offset if the ref's owner chain is a SliceTmemOp (optionally behind a TmemLayoutCastOp). Any other op producing the TMEM ref (arbitrary layout casts, allocs, other dialect ops) makes offset computation impossible and NotImplementedError is raised.","triggerScenarios":"Aliasing a TMEM ref whose owner is not slice_tmem / tmem_layout_cast(slice_tmem) — e.g. aliasing a raw TMEM allocation or the result of another TMEM-transforming op.","commonSituations":"Blackwell tcgen05 kernels chaining multiple TMEM layout operations before an aliased view; evolving JAX versions adding new TMEM ops the alias path doesn't recognize.","solutions":["Ensure the aliased TMEM ref comes directly from a tmem slice op (slice before any other transform)","Apply the alias view earlier in the op chain, immediately after slicing","Simplify/reorder TMEM layout casts so a slice_tmem is the immediate (or one-hop) owner","Upgrade JAX / report upstream with a minimal kernel repro"],"exampleFix":"# before\nr = tmem_alloc(...)            # owner not a slice op\nv = r.view(dtype)\n# after\nr = tmem_alloc(...)\nr = r[0:n]                     # slice_tmem owner\nv = r.view(dtype)","handlingStrategy":"fallback","validationCode":"owner = getattr(ref, 'owner', None)\nok = type(owner).__name__ == 'SliceTmemOp' or (\n     type(owner).__name__ == 'TmemLayoutCastOp' and\n     type(getattr(owner.operands[0], 'owner', None)).__name__ == 'SliceTmemOp')\nassert ok, 'TMEM alias must come from slice_tmem (optionally behind tmem_layout_cast)'","typeGuard":"def tmem_ref_aliasable(ref) -> bool:\n    o = getattr(ref, 'owner', None)\n    if type(o).__name__ == 'SliceTmemOp':\n        return True\n    if type(o).__name__ == 'TmemLayoutCastOp':\n        return type(getattr(o.operands[0], 'owner', None)).__name__ == 'SliceTmemOp'\n    return False\n","tryCatchPattern":null,"preventionTips":["Slice TMEM allocations before applying views or layout casts","Minimize TMEM op chains feeding aliased refs","Add regression tests when upgrading JAX on tcgen05 kernels"],"tags":["jax","pallas","tmem","tcgen05","aliasing","not-implemented"],"backgroundTag":"unsupported-lowering-path","analyzedSha":"1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb","analyzedAt":"2026-08-27T09:53:25.647Z","schemaVersion":2},"datasetVersion":"2026-08-27T13:17:12.746Z"}