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

  1. 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
  2. Convert the array to a layout whose registers are 32/64-bit wide before store_tiled_async
  3. 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

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


AI-assisted analysis of jax-ml/jax@1e1c6a8fc0 (2026-08-27). Data as JSON: /api/errors/f92d0a7b5aa32221. Report an issue: GitHub.