{"record":{"id":"d830c104a4d7e464","repo":"jax-ml/jax","slug":"a-sparse-metadata-dtype-mismatch-expected-i2-got","errorCode":null,"errorMessage":"A sparse metadata dtype mismatch: expected i2, got {a_sparse_metadata.dtype}","messagePattern":"A sparse metadata dtype mismatch: expected i2, got (.+?)","errorType":"validation","errorClass":"ValueError","httpStatus":null,"severity":"error","filePath":"jax/experimental/mosaic/gpu/tcgen05.py","lineNumber":506,"sourceCode":"          f\"B scale shape[0] must be a multiple of 128 and >= N={n * num_cta},\"\n          f\" got {b_scale.shape[0]}\"\n      )\n    if b_scale.shape[1] != k_scales:\n      raise ValueError(\n          f\"B scale shape mismatch: expected ({b_scale.shape[0]}, {k_scales}),\"\n          f\" got {b_scale.shape}\"\n      )\n  if is_sparse:\n    sparse_group_elems = 8 if utils.bitwidth(a_element_type) == 4 else 4\n    # Each sparse group has 2 entries.\n    expected_meta_k = k // sparse_group_elems * 2\n    if a_sparse_metadata.shape != (m, expected_meta_k):\n      raise ValueError(\n          f\"A sparse metadata shape mismatch: expected {(m, expected_meta_k)},\"\n          f\" got {a_sparse_metadata.shape}\"\n      )\n    if a_sparse_metadata.dtype != ir.IntegerType.get_signless(2):\n      raise ValueError(\n          \"A sparse metadata dtype mismatch: expected i2, got\"\n          f\" {a_sparse_metadata.dtype}\"\n      )\n\n  # Step 3. Compute the operand descriptors.\n  if not isinstance(a, TMEMRef):\n    # Both dense and sparse matmul consume A with a K bytewidth of 32, only\n    # the group size is halved when it's sparse.\n    (\n        (a_desc_base, a_k_instr_strides),\n        (a_m_group_stride, a_k_group_stride),\n        a_fastest,\n    ) = mma_utils.create_descriptor(\n        a,\n        swizzle=a_swizzle,\n        group_size=(m_group_elems, k_group_elems // (1 + is_sparse)),\n        logical_k_major=False,\n        mma_bytewidth_k=32,","sourceCodeStart":488,"sourceCodeEnd":524,"githubUrl":"https://github.com/jax-ml/jax/blob/1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb/jax/experimental/mosaic/gpu/tcgen05.py#L488-L524","documentation":"Sparse MMA requires the A sparse metadata tensor to be i2 (signless 2-bit integer), matching the tcgen05 sparse metadata encoding. Any other dtype raises this error.","triggerScenarios":"Passing a_sparse_metadata with dtype i8, i32, or signed i2 variants instead of ir.IntegerType.get_signless(2).","commonSituations":"Metadata produced by JAX/torch as uint8 or int32 without repacking into packed i2; loading metadata from files with default integer dtypes.","solutions":["Pack the metadata into 2-bit signless integers (4 values per byte) and view the buffer as i2","Check a_sparse_metadata.dtype == ir.IntegerType.get_signless(2) before calling mma"],"exampleFix":"# before\nmeta = memref (m, meta_k) of i8\ntcgen05.mma(..., a_sparse_metadata=meta)\n# after\nmeta_i2 = packed_i2_metadata(m, meta_k)  # ir.IntegerType.get_signless(2)\ntcgen05.mma(..., a_sparse_metadata=meta_i2)","handlingStrategy":"validation","validationCode":"assert a_sparse_metadata.dtype == ir.IntegerType.get_signless(2)","typeGuard":"def is_i2_meta(t) -> bool:\n    return t.dtype == ir.IntegerType.get_signless(2)","tryCatchPattern":null,"preventionTips":["Pack 2-bit metadata explicitly before kernel launch","Assert dtype on every load of metadata buffers"],"tags":["gpu","mosaic","tcgen05","sparsity","dtype","metadata"],"backgroundTag":"dtype-mismatch","analyzedSha":"1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb","analyzedAt":"2026-08-27T09:53:25.647Z","schemaVersion":2},"datasetVersion":"2026-08-27T13:17:12.746Z"}