{"record":{"id":"c167090dda3d16f7","repo":"jax-ml/jax","slug":"unsupported-operand-type-element-type","errorCode":null,"errorMessage":"Unsupported operand type: {element_type}","messagePattern":"Unsupported operand type: (.+?)","errorType":"exception","errorClass":"NotImplementedError","httpStatus":null,"severity":"error","filePath":"jax/experimental/mosaic/gpu/mma.py","lineNumber":211,"sourceCode":"  if n != n2:\n    raise ValueError(f\"N mismatch: {n} != {n2}\")\n  if k != k2:\n    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:","sourceCodeStart":193,"sourceCodeEnd":229,"githubUrl":"https://github.com/jax-ml/jax/blob/1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb/jax/experimental/mosaic/gpu/mma.py#L193-L229","documentation":"Mosaic's mma supports only bf16, f16, f8e4m3fn, f8e5m2, i8 and i4 operand types. Any other element type raises NotImplementedError.","triggerScenarios":"Calling mma() with f64, f32, or an exotic integer width (e.g. i16) as operand element type.","commonSituations":"Porting kernels that assume tf32/f32 tensor cores; older GPUs or expecting full float32 MMA support.","solutions":["Downcast operands to a supported dtype (bf16/f16/f8/int8) before mma","Use plain elementwise multiply-accumulate for unsupported high-precision types","Keep an f32 accumulator with f16/bf16 operands for range"],"exampleFix":"// before\nacc = mma.mma(a_f32, b_f32, acc)\n// after\nacc = mma.mma(a_f32.astype(jnp.bfloat16), b_f32.astype(jnp.bfloat16), acc)","handlingStrategy":"validation","validationCode":"supported = {'bf16','f16','f8e4m3fn','f8e5m2','i8','i4'}\nassert str(a.mlir_dtype) in supported","typeGuard":null,"tryCatchPattern":null,"preventionTips":["Downcast to bf16/f16/int8 before mma; keep acc f32 for floats"],"tags":["jax","mosaic","mma","unsupported-dtype"],"backgroundTag":"unsupported-dtype","analyzedSha":"1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb","analyzedAt":"2026-08-27T09:53:25.647Z","schemaVersion":2},"datasetVersion":"2026-08-27T13:17:12.746Z"}