{"record":{"id":"9fa8e79227598773","repo":"jax-ml/jax","slug":"m-warps-must-be-1-2-or-4-but-got-m-warps","errorCode":null,"errorMessage":"m_warps must be 1, 2, or 4, but got {m_warps=}","messagePattern":"m_warps must be 1, 2, or 4, but got (.+?)","errorType":"exception","errorClass":"ValueError","httpStatus":null,"severity":"error","filePath":"jax/experimental/mosaic/gpu/mma.py","lineNumber":37,"sourceCode":"from jax.experimental.mosaic.gpu import fragmented_array as fa\nfrom jaxlib.mlir import ir\nfrom jaxlib.mlir.dialects import llvm\nfrom jaxlib.mlir.dialects import vector\nimport numpy as np\nfrom . import utils\n\n\nSUPPORTED_F8_TYPES = (ir.Float8E4M3FNType, ir.Float8E5M2Type)\n\n\nclass MMALayouts:\n  \"\"\"Container for MMA layouts, providing a convenient way to create\n  layouts for MMA operands based on warp configuration.\n  \"\"\"\n\n  def __init__(self, element_type: ir.Type | typing.DTypeLike, *, m_warps: int = 4):\n    if m_warps not in (1, 2, 4):\n      raise ValueError(f\"m_warps must be 1, 2, or 4, but got {m_warps=}\")\n    n_warps = 4 // m_warps\n    if isinstance(element_type, ir.Type):\n      bitwidth = utils.bitwidth(element_type)\n    else:\n      bitwidth = dtypes.itemsize_bits(element_type)\n    elems_per_reg = 32 // bitwidth\n    k = 8 * elems_per_reg\n    sub_k = 4 * elems_per_reg\n    self.lhs = fa.TiledLayout(\n        fa.Tiling(((m_warps * 16, k), (16, sub_k), (8, sub_k), (elems_per_reg,))),\n        warp_dims=(-7, fa.Replicated(4 // m_warps)),\n        lane_dims=(-3, -2),\n        vector_dim=-1,\n        _check_canonical=False,\n    ).canonicalize()\n    self.rhs = fa.TiledLayout(\n        fa.Tiling(((k, n_warps * 8), (sub_k, 8), (elems_per_reg, 1))),\n        warp_dims=(fa.Replicated(4 // n_warps), -5,),","sourceCodeStart":19,"sourceCodeEnd":55,"githubUrl":"https://github.com/jax-ml/jax/blob/1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb/jax/experimental/mosaic/gpu/mma.py#L19-L55","documentation":"FragmentedArray MMA layout construction partitions 4 warps between M and N dimensions; m_warps must therefore divide into 4, and only 1, 2, or 4 are valid. Other values raise ValueError.","triggerScenarios":"Passing m_warps=3, 8, 0, etc. to an MMA layout constructor / FragmentedArray helper that takes m_warps.","commonSituations":"Parameterizing kernel warp splits from config strings or sweep scripts without validating against {1,2,4}.","solutions":["Use m_warps in (1, 2, 4)","Validate user-supplied warp configs before passing them in","If you need more warps, scale via repetition/other tiling, not m_warps"],"exampleFix":"// before\nlayout = mma.MMALayout(dt, m_warps=3)\n// after\nlayout = mma.MMALayout(dt, m_warps=2)","handlingStrategy":"validation","validationCode":"assert m_warps in (1, 2, 4), 'm_warps must divide 4'","typeGuard":null,"tryCatchPattern":null,"preventionTips":["Validate sweep configs against {1,2,4} before kernel construction"],"tags":["jax","mosaic","mma","warps","validation"],"backgroundTag":"invalid-configuration-value","analyzedSha":"1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb","analyzedAt":"2026-08-27T09:53:25.647Z","schemaVersion":2},"datasetVersion":"2026-08-27T13:17:12.746Z"}