{"record":{"id":"c0abd9319e983834","repo":"jax-ml/jax","slug":"only-f16-wgmma-supports-transposes","errorCode":null,"errorMessage":"Only f16 WGMMA supports transposes","messagePattern":"Only f16 WGMMA supports transposes","errorType":"validation","errorClass":"ValueError","httpStatus":null,"severity":"error","filePath":"jax/experimental/mosaic/gpu/wgmma.py","lineNumber":154,"sourceCode":"  if not _supported_wgmma_types(out_ty, b_element_type):\n    raise ValueError(f\"Unsupported wgmma types {(out_ty, b_element_type)=}\")\n  if n % 8:\n    raise ValueError\n\n  bf16 = ir.BF16Type.get()\n  f16 = ir.F16Type.get()\n  i8 = ir.IntegerType.get_signless(8)\n  i32 = ir.IntegerType.get_signless(32)\n  i64 = ir.IntegerType.get_signless(64)\n  f8e5m2 = ir.Float8E5M2Type.get()\n  f8e4m3fn = ir.Float8E4M3FNType.get()\n  if b_k_stride % 16:\n    raise ValueError\n  assert bytewidth(a_element_type) == bytewidth(b_element_type)\n  # Only 16-bit types support transposes\n  supports_transpose = bytewidth(b_element_type) == 2\n  if not supports_transpose and (a_transpose or b_transpose):\n    raise ValueError(\"Only f16 WGMMA supports transposes\")\n  if a_in_regs := isinstance(a, fa.FragmentedArray):\n    if a.mlir_dtype not in {bf16, f16, i8, f8e5m2, f8e4m3fn}:\n      raise ValueError(f\"Unsupported A register array dtype: {a.mlir_dtype}\")\n    # Column count must be equal to swizzle // bytewidth.\n    elt_bytewidth = utils.bytewidth(a_element_type)\n    swizzle_elems = swizzle // elt_bytewidth\n    if a.shape != (64, swizzle_elems):\n      raise ValueError(\"Unsupported A register array shape\")\n    if a.layout not in {fa.WGMMA_LAYOUT, fa.WGMMA_LAYOUT_8BIT}:\n      raise ValueError(\"Unsupported A register array layout\")\n    if a_k_stride is not None or a_transpose is not None:\n      raise ValueError(\"Unsupported WGMMA features with A in registers\")\n  else:\n    if a_k_stride is None or a_k_stride % 16:\n      raise ValueError\n    if a_transpose is None:\n      raise ValueError\n","sourceCodeStart":136,"sourceCodeEnd":172,"githubUrl":"https://github.com/jax-ml/jax/blob/1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb/jax/experimental/mosaic/gpu/wgmma.py#L136-L172","documentation":"Raised at wgmma.py:154 when a_transpose or b_transpose is requested but the operand bytewidth is not 2 — the hardware only supports transposed WGMMA operands for 16-bit (f16/bf16) types.","triggerScenarios":"Calling wgmma.wgmma(..., a_transpose=True or b_transpose=True) with i8/f8/s32 operands.","commonSituations":"Porting an f16 attention kernel to int8 quantized operands and keeping the transpose flags; enabling transpose on low-precision B stored in SMEM.","solutions":["Remove a_transpose/b_transpose for non-16-bit dtypes and physically transpose the data instead (swap index math or pre-transpose in SMEM)","Convert operands to f16/bf16 if transposes are essential","Use transpose via TMA layout rather than the wgmma flags"],"exampleFix":"# before\nacc = wgmma.wgmma(a_i8, b_i8, acc, b_transpose=True)\n# after\nb_t = utils.transpose_smem(b)  # or load B pre-transposed\nacc = wgmma.wgmma(a_i8, b_t, acc)","handlingStrategy":"validation","validationCode":"if bytewidth(b_element_type) != 2:\n    assert not a_transpose and not b_transpose, 'transpose requires 16-bit operands'","typeGuard":null,"tryCatchPattern":null,"preventionTips":["Pre-transpose int8 data in memory instead of using transpose flags","Reserve wgmma transpose flags for f16/bf16 kernels"],"tags":["jax","mosaic-gpu","wgmma","transpose","dtype"],"backgroundTag":"unsupported-type-combination","analyzedSha":"1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb","analyzedAt":"2026-08-27T09:53:25.647Z","schemaVersion":2},"datasetVersion":"2026-08-27T13:17:12.746Z"}