jax-ml/jax · error · NotImplementedError

A scale address calculation for multiple M tiles

Error message

A scale address calculation for multiple M tiles

What it means

Block-scale TMEM addressing for the A scale tensor across multiple M tiles is not implemented; when both a_scale and b_scale are given, mma requires m_groups == 1.

Source

Thrown at jax/experimental/mosaic/gpu/tcgen05.py:599

      a_mk = a.slice(slice(None), utils.ds(ki * a_k_group_elems, a_k_group_elems)).address
    else:
      assert a_desc_base is not None
      a_offset = mi * a_m_group_stride + ki * a_k_group_stride
      a_mk = (a_desc_base[0], a_desc_base[1] + mma_utils.encode_addr(a_offset))
    b_offset = ni * b_n_group_stride + ki * b_k_group_stride
    b_nk = (b_desc_base[0], b_desc_base[1] + mma_utils.encode_addr(b_offset))
    if a_sparse_addr_base is not None:
      if n_groups != 1 or m_groups != 1:
        raise NotImplementedError("A sparse metadata address calculation for multiple tiles")
      sparse_group_elems = 8 if utils.bitwidth(mma_a_element_type) == 4 else 4
      # Each sparse group has 2 entries, each TMEM column holds 16 i2 entries.
      cols_per_k_group = k_group_elems // sparse_group_elems * 2 // 16
      a_sparse_addr = arith.addi(a_sparse_addr_base, utils.c(ki * cols_per_k_group, i32))
    else:
      a_sparse_addr = None
    if a_scale_addr_base is not None and b_scale_addr_base is not None:
      if m_groups != 1:
        raise NotImplementedError("A scale address calculation for multiple M tiles")
      if n_groups != 1:
        raise NotImplementedError("B scale address calculation for multiple N tiles")
      assert scale_block is not None  # For type checkers.
      assert k_group_elems % (scale_block * 4) == 0
      assert m_group_elems % 32 == 0 and n_group_elems % (8 * num_cta) == 0
      k_scales_per_group = k_group_elems // (scale_block * 4)
      a_scale_addr = arith.addi(
          a_scale_addr_base,
          utils.c(ki * k_scales_per_group * a_scale_m_stride, i32),
      )
      b_scale_addr = arith.addi(
          b_scale_addr_base,
          utils.c(ki * k_scales_per_group * b_scale_n_stride, i32)
      )
    else:
      a_scale_addr = b_scale_addr = None
    acc = accumulate if ki == 0 else true
    ni_lane_group, ni_col = ni // n_col_groups, ni % n_col_groups

View on GitHub (pinned to 1e1c6a8fc0)

Solutions

  1. Split the M dimension manually and issue one mma per M tile
  2. Ensure the per-call m equals the full M so m_groups==1

Example fix

# before
tcgen05.mma(a, b, d, a_scale=asc, b_scale=bsc, m=256)  # m_groups=2
# after
for mi in range(2):
  tcgen05.mma(a.slice(mi), b, d.slice(mi), a_scale=asc, b_scale=bsc, m=128)
Defensive patterns

Strategy: validation

Validate before calling

assert m_groups == 1  # block-scaled A scale addressing supports single M tile

Prevention

When it happens

Trigger: Block-scaled mma() with a_scale/b_scale supplied and a tile loop where m_groups > 1.

Common situations: Large-M MXFP8/FP4 GEMM where M exceeds one instruction tile (128) and the code relies on mma's internal tiling.

Understand the failure class

Background: UnsupportedOperationException and "is not supported" errors: when a library deliberately refuses a call — this error's family across 30 libraries.

Related errors


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