{"record":{"id":"9dd5cf6790db2adb","repo":"jax-ml/jax","slug":"only-s32-accumulator-supported-for-integer-operand","errorCode":null,"errorMessage":"Only s32 accumulator supported for integer operands.","messagePattern":"Only s32 accumulator supported for integer operands\\.","errorType":"exception","errorClass":"NotImplementedError","httpStatus":null,"severity":"error","filePath":"jax/experimental/mosaic/gpu/mma.py","lineNumber":214,"sourceCode":"    raise ValueError(f\"K mismatch: {k} != {k2}\")\n\n  # todo(cperivol): A tile shape can have dimensions that are higher\n  # multiples of the mma op size as long as those dimensions are not\n  # sharded across warps.\n  i4 = ir.IntegerType.get_signless(4)\n  i8 = ir.IntegerType.get_signless(8)\n  i32 = ir.IntegerType.get_signless(32)\n  bf16 = ir.BF16Type.get()\n  f16 = ir.F16Type.get()\n  f8e4m3fn = ir.Float8E4M3FNType.get()\n  f8e5m2 = ir.Float8E5M2Type.get()\n  if (element_type := a.mlir_dtype) != b.mlir_dtype:\n    raise ValueError(f\"Dtype mismatch: {a.mlir_dtype} != {b.mlir_dtype}\")\n  if element_type not in (bf16, f16, f8e4m3fn, f8e5m2, i8, i4):\n    raise NotImplementedError(f\"Unsupported operand type: {element_type}\")\n  if isinstance(element_type, ir.IntegerType):\n    if acc.mlir_dtype != i32:\n      raise NotImplementedError(\"Only s32 accumulator supported for integer operands.\")\n    if not acc.is_signed:\n      raise ValueError(\"Only signed accumulator supported for integer operands.\")\n  elif acc.mlir_dtype != ir.F32Type.get():\n    raise NotImplementedError(\"Only f32 accumulator supported for floating operands.\")\n\n  can_infer_from_acc_layout = (\n      isinstance(acc.layout, fa.TiledLayout)\n      and len(acc.layout.base_tile_shape) == 2\n      and acc.layout.base_tile_shape[0] % 16 == 0\n  )\n  if not can_infer_from_acc_layout:\n    raise ValueError(\"Expected MMALayouts.acc for acc\")\n  m_warps = acc.layout.base_tile_shape[0] // 16  # type: ignore\n  layouts = MMALayouts(element_type, m_warps=m_warps)\n  if layouts.lhs != a.layout:\n    raise ValueError(\"Expected MMALayouts.lhs layout for A\")\n  if layouts.rhs != b.layout:\n    raise ValueError(\"Expected MMALayouts.rhs layout for B\")","sourceCodeStart":196,"sourceCodeEnd":232,"githubUrl":"https://github.com/jax-ml/jax/blob/1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb/jax/experimental/mosaic/gpu/mma.py#L196-L232","documentation":"For integer MMA operands (i8/i4), the accumulator must be 32-bit integers (s32); PTX integer MMA only produces i32 results. Any other accumulator dtype raises NotImplementedError.","triggerScenarios":"Calling mma() with i8 operands but an f32 or i8 accumulator.","commonSituations":"Copy-pasting an fp16 kernel's f32 accumulator setup into an int8 quantized kernel.","solutions":["Allocate the accumulator as jnp.int32 (signed) for integer operands","Keep f32 accumulators only for floating-point operand types"],"exampleFix":"// before\nacc = fa.from_tensor(jnp.zeros((m, n), jnp.float32))\nacc = mma.mma(a_i8, b_i8, acc)\n// after\nacc = fa.from_tensor(jnp.zeros((m, n), jnp.int32))\nacc = mma.mma(a_i8, b_i8, acc)","handlingStrategy":"validation","validationCode":"if isinstance(a.mlir_dtype, ir.IntegerType):\n    assert acc.mlir_dtype == ir.IntegerType.get_signless(32)","typeGuard":null,"tryCatchPattern":null,"preventionTips":["Pair int8 operands with a jnp.int32 accumulator"],"tags":["jax","mosaic","mma","accumulator","int8"],"backgroundTag":"accumulator-dtype-mismatch","analyzedSha":"1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb","analyzedAt":"2026-08-27T09:53:25.647Z","schemaVersion":2},"datasetVersion":"2026-08-27T13:17:12.746Z"}