{"record":{"id":"1174bc2108b7cd82","repo":"jax-ml/jax","slug":"only-wgmma-layouts-supported-in-wgmmaaccumulator","errorCode":null,"errorMessage":"Only WGMMA layouts supported in WGMMAAccumulator","messagePattern":"Only WGMMA layouts supported in WGMMAAccumulator","errorType":"validation","errorClass":"ValueError","httpStatus":null,"severity":"error","filePath":"jax/experimental/mosaic/gpu/wgmma.py","lineNumber":90,"sourceCode":"      raise TypeError(\"PTX does not support unsigned WGMMA accumulators\")\n    f32 = ir.F32Type.get()\n    if dtype is None:\n      dtype = f32\n    if isinstance(dtype, ir.IntegerType):\n      zero = arith.constant(dtype, ir.IntegerAttr.get(dtype, 0))\n    else:\n      zero = arith.constant(dtype, ir.FloatAttr.get(dtype, 0.0))\n    return cls.from_registers(\n        fa.FragmentedArray.splat(\n            zero, (m, n), fa.WGMMA_LAYOUT, is_signed=is_signed\n        )\n    )\n\n  @classmethod\n  def from_registers(cls, registers, sync=True):\n    original_layout = registers.layout\n    if registers.layout != fa.WGMMA_LAYOUT and registers.layout != fa.WGMMA_LAYOUT_ACC_32BIT:\n      raise ValueError(\"Only WGMMA layouts supported in WGMMAAccumulator\")\n    if utils.bitwidth(registers.mlir_dtype) == 32:\n      registers = registers.to_layout(fa.WGMMA_LAYOUT_ACC_32BIT)\n    return cls(_value=registers, _original_layout=original_layout, _sync=sync)\n\n  def tree_flatten(self):\n    return (self._value,), (self._original_layout,)\n\n  @classmethod\n  def tree_unflatten(cls, aux, value):\n    return cls(_value=value[0], _original_layout=aux[0], _sync=False)\n\n\ndef _supported_wgmma_types(dtype, abtype) -> bool:\n  input_types_are = lambda ty: isinstance(abtype, ty)\n  f16_acc_types = (ir.F16Type, ir.Float8E5M2Type, ir.Float8E4M3FNType)\n  if isinstance(dtype, ir.F32Type):\n    return any(input_types_are(ty) for ty in (ir.FloatTF32Type, ir.BF16Type, *f16_acc_types))\n  elif isinstance(dtype, ir.F16Type):","sourceCodeStart":72,"sourceCodeEnd":108,"githubUrl":"https://github.com/jax-ml/jax/blob/1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb/jax/experimental/mosaic/gpu/wgmma.py#L72-L108","documentation":"WGMMAAccumulator.from_registers (wgmma.py:90) only accepts FragmentedArrays laid out as fa.WGMMA_LAYOUT or fa.WGMMA_LAYOUT_ACC_32BIT, because the accumulator must present registers exactly as the wgmma instruction produces them.","triggerScenarios":"Calling wgmma.WGMMAAccumulator.from_registers(regs) where regs is a FragmentedArray in a row-major/blocked/register-tensor layout instead of a WGMMA layout; also hit via lowering rules (_mgpu_wgmma_op_lowering_rule, accumulator store lowering, kv_loop in FlashAttention examples).","commonSituations":"Using wgmma(...) output that was already relayouted (e.g. after arithmetic in a different layout); passing a tensor created via fa.make_tensor and forgetting .to_layout(fa.WGMMA_LAYOUT_ACC_32BIT); version changes that renamed layouts.","solutions":["Convert registers first: regs.to_layout(fa.WGMMA_LAYOUT_ACC_32BIT if bitwidth==32 else fa.WGMMA_LAYOUT)","Ensure you pass the direct result of wgmma.wgmma (or wgmma_m64) without intermediate relayouts","Check fa.WGMMA_LAYOUT constants exist in your JAX version; upgrade or use the renamed constants"],"exampleFix":"# before\nacc = wgmma.WGMMAAccumulator.from_registers(regs)  # regs in row-major\n# after\nregs = regs.to_layout(fa.WGMMA_LAYOUT_ACC_32BIT)\nacc = wgmma.WGMMAAccumulator.from_registers(regs)","handlingStrategy":"validation","validationCode":"assert regs.layout in (fa.WGMMA_LAYOUT, fa.WGMMA_LAYOUT_ACC_32BIT), f'bad layout {regs.layout}'","typeGuard":"def is_wgmma_layout(a): return a.layout in (fa.WGMMA_LAYOUT, fa.WGMMA_LAYOUT_ACC_32BIT)","tryCatchPattern":null,"preventionTips":["Pass wgmma results straight into the accumulator","Relayout explicitly to WGMMA_LAYOUT_ACC_32BIT for 32-bit dtypes"],"tags":["jax","mosaic-gpu","wgmma","fragmented-array","layout"],"backgroundTag":"unsupported-layout-conversion","analyzedSha":"1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb","analyzedAt":"2026-08-27T09:53:25.647Z","schemaVersion":2},"datasetVersion":"2026-08-27T13:17:12.746Z"}