{"record":{"id":"ed3ef3ce3c5d0f51","repo":"jax-ml/jax","slug":"only-32-bit-types-supported","errorCode":null,"errorMessage":"Only 32-bit types supported","messagePattern":"Only 32-bit types supported","errorType":"validation","errorClass":"ValueError","httpStatus":null,"severity":"error","filePath":"jax/experimental/mosaic/gpu/utils.py","lineNumber":2021,"sourceCode":"  if (x_bitwidth := bitwidth(result_type)) < 32:\n    bits_ty = ir.IntegerType.get_signless(x_bitwidth)\n    y_vec = bitcast(y, ir.VectorType.get((32 // x_bitwidth,), bits_ty))\n    y = vector.extract(\n        y_vec,\n        dynamic_position=[],\n        static_position=ir.DenseI64ArrayAttr.get([0]),\n    )\n  return bitcast(y, result_type)\n\n\nReductionKind = nvvm.ReductionKind\n\n\ndef redux(x: ir.Value, mask: ir.Value, kind: ReductionKind):\n  i32 = ir.IntegerType.get_signless(32)\n  if isinstance(vec_ty := x.type, ir.VectorType):\n    if bitwidth(vec_ty.element_type) != 32:\n      raise ValueError(\"Only 32-bit types supported\")\n    [vec_len] = vec_ty.shape\n    result = llvm.mlir_undef(x.type)\n    for i in range(vec_len):\n      xi = llvm.extractelement(x, arith.constant(i32, i))\n      yi = redux(xi, mask, kind)\n      result = llvm.insertelement(result, yi, arith.constant(i32, i))\n    return result\n  if bitwidth(x.type) != 32:\n    raise ValueError(\"Only 32-bit scalar types supported\")\n  if isinstance(x.type, ir.IntegerType):\n    pass\n  elif isinstance(x.type, ir.F32Type):\n    if get_arch().major != 10:\n      raise ValueError(\"F32 redux only supported on Blackwell GPUs\")\n  else:\n    raise NotImplementedError(x.type)\n  assert mask.type == i32\n  extra_kwargs: dict[str, Any] = {}","sourceCodeStart":2003,"sourceCodeEnd":2039,"githubUrl":"https://github.com/jax-ml/jax/blob/1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb/jax/experimental/mosaic/gpu/utils.py#L2003-L2039","documentation":"redux lowers to PTX redux.sync, which on NVIDIA hardware only operates on 32-bit values. When x is a vector, redux recursively extracts elements and checks each is 32-bit; a vector of f16/bf16/i8/i64 elements fails this check.","triggerScenarios":"Calling redux(x, mask, kind) where x is e.g. vector<4xbf16> or vector<8xi8> — any vector whose element bitwidth != 32.","commonSituations":"Feeding low-precision accumulators (bf16/f16) from attention or GEMM epilogues directly into redux without upcasting; assuming redux works like a generic reduce utility across dtypes.","solutions":["Upcast the vector to a 32-bit element type before redux (e.g. via arith.extf to f32 or fp_ext) and downcast after","Use a different reduction strategy for sub-32-bit types (shuffle-based warp_reduce or LLVM vector reduce ops)","Check dtype at the kernel boundary and insert explicit conversion nodes"],"exampleFix":"# before\nresult = redux(x_bf16, mask, ReductionKind.ADD)\n# after\nx_f32 = arith.extf(f32, x_bf16)\nresult = redux(x_f32, mask, ReductionKind.ADD)","handlingStrategy":"type-guard","validationCode":"def is_vec32(v):\n    t = ir.VectorType(v.type)\n    return bitwidth(t.element_type) == 32\nassert is_vec32(x), 'upcast vector to 32-bit elements before redux'","typeGuard":"def is_redux_capable_vector(v) -> bool:\n    return isinstance(v.type, ir.VectorType) and bitwidth(ir.VectorType(v.type).element_type) == 32","tryCatchPattern":null,"preventionTips":["Upcast f16/bf16 vectors to f32 at reduction boundaries","Standardize accumulator dtypes as 32-bit in kernel style guides"],"tags":["mosaic-gpu","redux","dtype","hardware-limit"],"backgroundTag":"unsupported-dtype","analyzedSha":"1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb","analyzedAt":"2026-08-27T09:53:25.647Z","schemaVersion":2},"datasetVersion":"2026-08-27T13:17:12.746Z"}