{"record":{"id":"f994d7d7af47d012","repo":"jax-ml/jax","slug":"sparse-metadata-format-not-implemented-for-operan","errorCode":null,"errorMessage":"Sparse metadata format not implemented for {operand_dtype=}","messagePattern":"Sparse metadata format not implemented for (.+?)","errorType":"validation","errorClass":"NotImplementedError","httpStatus":null,"severity":"error","filePath":"jax/_src/pallas/mosaic_gpu/helpers.py","lineNumber":220,"sourceCode":"  if meta.ndim != 3:\n    raise ValueError(\n        \"Expected metadata to be 3-dimensional (M, K // 4, 2), but it is\"\n        f\" {meta.ndim}D\"\n    )\n  m, k, _2 = meta.shape\n  if _2 != 2:\n    raise ValueError(\n        \"Expected the trailing dimension of the metadata to be 2, got:\"\n        f\" {meta.shape[-1]}\"\n    )\n  k *= 2\n  bitsize = dtypes.itemsize_bits(operand_dtype)\n  if bitsize == 8:\n    meta_tiled = meta.reshape(m // 128, 128, k // 64, 64).transpose(0, 2, 1, 3)\n  elif bitsize == 16:\n    meta_tiled = meta.reshape(m // 128, 8, 2, 8, k // 64, 4, 2, 8).transpose(0, 4, 1, 6, 3, 5, 2, 7)\n  else:\n    raise NotImplementedError(\n        f\"Sparse metadata format not implemented for {operand_dtype=}\"\n    )\n  return meta_tiled.reshape(m // 128, k // 64, 128, 64)\n\n\ndef find_swizzle(minor_dim_bits: int, what: str = \"\"):\n  \"\"\"Returns the largest swizzle that can be applied to a memory region.\n\n  Swizzling is usually necessary when dealing with 2D data in SMEM, especially\n  if the reference is used as an MMA operand. The returned swizzle is usually\n  applied as ``plgpu`` transform:\n\n    transforms = (\n        plgpu.TilingTransform((8, 8 * swizzle // elem_bits)),\n        plgpu.SwizzleTransform(swizzle))\n    )\n\n  Args:","sourceCodeStart":202,"sourceCodeEnd":238,"githubUrl":"https://github.com/jax-ml/jax/blob/1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb/jax/_src/pallas/mosaic_gpu/helpers.py#L202-L238","documentation":"format_tcgen05_sparse_metadata implements tiling patterns only for 8-bit and 16-bit operand dtypes. Other bit sizes (e.g. 32-bit floats/ints) have no tcgen05 sparse metadata packing defined, so NotImplementedError is raised.","triggerScenarios":"Calling format_tcgen05_sparse_metadata(meta, operand_dtype) with operand_dtype of itemsize other than 8 or 16 bits, e.g. jnp.float32 or jnp.bfloat16 is 16-bit ok but jnp.int32 is not.","commonSituations":"Attempting sparse matmul with fp32 accumulators passed as the operand dtype by mistake; exploring sparsity with dtypes the hardware sparse path doesn't support.","solutions":["Use an 8- or 16-bit operand dtype such as float8_e4m3fn, float8_e5m2, float16, or bfloat16","Pass the actual MMA operand dtype, not the accumulator dtype (accumulators are usually fp32)","If you need wider dtypes, use the dense (non-sparse) path"],"exampleFix":"// before\nmeta_tiled = format_tcgen05_sparse_metadata(meta, jnp.float32)\n\n// after\nmeta_tiled = format_tcgen05_sparse_metadata(meta, jnp.float8_e4m3fn)","handlingStrategy":"type-guard","validationCode":"from jax import dtypes\nbits = dtypes.itemsize_bits(operand_dtype)\nassert bits in (8, 16), f'sparse metadata unsupported for {bits}-bit operands'","typeGuard":"def sparse_dtype_supported(operand_dtype) -> bool:\n    from jax import dtypes\n    return dtypes.itemsize_bits(operand_dtype) in (8, 16)","tryCatchPattern":"try:\n    return format_tcgen05_sparse_metadata(meta, operand_dtype)\nexcept NotImplementedError:\n    # fall back to dense operand dtype or dense matmul\n    return format_tcgen05_sparse_metadata(meta, jnp.float8_e4m3fn)","preventionTips":["Use 8/16-bit MMA operand dtypes for the sparse path","Pass the operand dtype, not the fp32 accumulator dtype","Branch to a dense kernel for wider dtypes"],"tags":["jax","pallas","mosaic-gpu","sparse","tcgen05","dtype","not-implemented"],"backgroundTag":"dtype-not-supported","analyzedSha":"1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb","analyzedAt":"2026-08-27T09:53:25.647Z","schemaVersion":2},"datasetVersion":"2026-08-27T13:17:12.746Z"}