jax-ml/jax · error · NotImplementedError
Unsupported register bitwidth: {reg_bitwidth}
Error message
Unsupported register bitwidth: {reg_bitwidth} What it means
To emit the inline st.async PTX instruction, store_tiled_async must bitcast each register to a 32- or 64-bit PTX register type (b32 with 'r' constraint or b64 with 'l'); any other register bitwidth (e.g. 8- or 128-bit registers) has no corresponding PTX encoding so it raises NotImplementedError.
Source
Thrown at jax/experimental/mosaic/gpu/fragmented_array.py:3850
reg_ty = ir.VectorType(reg.type)
element_bitwidth = utils.bitwidth(reg_ty.element_type)
if (
isinstance(reg_ty.element_type, ir.FloatType)
and element_bitwidth <= 8
):
narrow_int = ir.IntegerType.get_signless(element_bitwidth)
reg = vector.bitcast(ir.VectorType.get(reg_ty.shape, narrow_int), reg)
reg_bitwidth = utils.bitwidth(reg_ty)
if reg_bitwidth == 32:
ptx_constraint = "r"
ptx_type = "b32"
reg = utils.bitcast(reg, i32)
elif reg_bitwidth == 64:
ptx_constraint = "l"
ptx_type = "b64"
reg = utils.bitcast(reg, i64)
else:
raise NotImplementedError(f"Unsupported register bitwidth: {reg_bitwidth}")
llvm.inline_asm(
ir.Type.parse("!llvm.void"),
[cluster_ptr, reg, cluster_barrier_ptr],
f"st.async.cluster.shared::cluster.mbarrier::complete_tx::bytes.{ptx_type} [$0], $1, [$2];",
f"l,{ptx_constraint},l",
has_side_effects=True,
)
def _store_register_atomic(
self,
base_ptr: ir.Value,
vreg: ir.Value,
atomic: Literal["add", "max", "min", "and", "or", "xor"],
is_smem: bool,
multimem: bool = False,
cluster_barrier_ptr: ir.Value | None = None,
):
i32 = ir.IntegerType.get_signless(32)View on GitHub (pinned to 1e1c6a8fc0)
Solutions
- Ensure sub-32-bit elements are packed to at least 32-bit registers (vec_len a multiple of 2 for 16-bit, 4 for 8-bit) via the layout's vec_size
- Convert the array to a layout whose registers are 32/64-bit wide before store_tiled_async
- Use the non-async store path for exotic widths
Example fix
// before fa8.store_tiled_async(ref) # i8 regs < 32 bits // after fa8 = fa8.to_layout(TiledLayout(vec_size=4)) fa8.store_tiled_async(ref)
Defensive patterns
Strategy: validation
Validate before calling
from jax.experimental.mosaic.gpu import utils
reg_bitwidth = utils.bitwidth(fa.mlir_dtype) * fa.layout.vec_size # approximate reg width
if reg_bitwidth not in (32, 64):
fa = fa.to_layout(TiledLayout(vec_size=max(1, 32 // utils.bitwidth(fa.mlir_dtype)))) Try / catch
try:
fa.store_tiled_async(ref, ...)
except NotImplementedError as e:
if 'register bitwidth' in str(e).lower():
fa.store_tiled(ref) # sync fallback
else:
raise Prevention
- Choose vec_size so elements pack into 32/64-bit registers
- Prefer 32-bit-friendly element types for async stores
When it happens
Trigger: Storing values whose fragment registers are not 32 or 64 bits wide — e.g. i8/f16 vector fragments that were not vectorized to 32-bit granularity, or unusual wide types.
Common situations: Async-storing sub-32-bit element types without ensuring vec_len pairs them into 32-bit registers; layouts with odd vec_size producing non-standard register widths.
Related errors
- Multimem refs are not supported in store_tiled_async
- Replicated dimensions are not supported
- f32 not supported for async atomics
- f32 only supports add atomics, got {atomic}
- f16/bf16 not supported for async atomics
AI-assisted analysis of jax-ml/jax@1e1c6a8fc0 (2026-08-27).
Data as JSON: /api/errors/f92d0a7b5aa32221.
Report an issue: GitHub.