{"record":{"id":"1540ea55b6a0f830","repo":"jax-ml/jax","slug":"expected-metadata-to-be-3-dimensional-m-k-4","errorCode":null,"errorMessage":"Expected metadata to be 3-dimensional (M, K // 4, 2), but it is {meta.ndim}D","messagePattern":"Expected metadata to be 3-dimensional \\(M, K // 4, 2\\), but it is (.+?)D","errorType":"validation","errorClass":"ValueError","httpStatus":null,"severity":"error","filePath":"jax/_src/pallas/mosaic_gpu/helpers.py","lineNumber":203,"sourceCode":"\ndef format_tcgen05_sparse_metadata(meta, operand_dtype):\n  \"\"\"Formats the sparse metadata for tcgen05.mma into the expected format.\n\n  See\n  https://docs.nvidia.com/cuda/parallel-thread-execution/#tcgen05-sparse-matrices-sparsity-selector-kind-f16-m128-256\n  for the documentation of the required layouts. The array can be copied into\n  SMEM, from where ``plgpu.async_copy_sparse_metadata_to_tmem`` can be used to\n  copy it over to TMEM. The formatting of the array depends on the data type of\n  the operands to the sparse MMA operation.\n\n  Args:\n    meta: Metadata of shape (M, K // 4, 2).\n    dtype: Data type of MMA operands.\n  \"\"\"\n  if meta.dtype != dtypes.uint2:\n    raise ValueError(f\"Expected metadata dtype to be uint2, got: {meta.dtype}\")\n  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=}\"","sourceCodeStart":185,"sourceCodeEnd":221,"githubUrl":"https://github.com/jax-ml/jax/blob/1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb/jax/_src/pallas/mosaic_gpu/helpers.py#L185-L221","documentation":"format_tcgen05_sparse_metadata expects metadata shaped (M, K // 4, 2) — exactly 3D. Higher- or lower-rank arrays are rejected because the subsequent reshape/transpose logic assumes this exact structure.","triggerScenarios":"Passing a 2D packed metadata array or a 4D batched array to format_tcgen05_sparse_metadata.","commonSituations":"Pre-packed/flattened metadata from a data pipeline; vmap/batching accidentally adding a leading dimension; misconverting the (M, K//4, 2) layout from documentation.","solutions":["Reshape to (M, K // 4, 2) before the call, e.g. meta.reshape(m, k // 4, 2) after ensuring dtype uint2","Remove accidental batch axes (squeeze) introduced by vmap or stacking"],"exampleFix":"// before\nmeta = packed.reshape(m, k // 2)  # 2D, wrong\n\n// after\nmeta = packed.reshape(m, k // 4, 2).astype(jnp.uint2)","handlingStrategy":"validation","validationCode":"assert meta.ndim == 3, f'expected (M, K//4, 2), got {meta.shape}'","typeGuard":"def metadata_shape_ok(meta) -> bool:\n    return meta.ndim == 3 and meta.shape[-1] == 2","tryCatchPattern":"try:\n    return format_tcgen05_sparse_metadata(meta, dt)\nexcept ValueError:\n    meta = meta.reshape(meta.shape[0], -1 // 2 if meta.ndim == 2 else meta.shape[1], 2)\n    return format_tcgen05_sparse_metadata(meta.astype(jnp.uint2), dt)","preventionTips":["Construct metadata as (M, K//4, 2) from the start","squeeze batch axes introduced by vmap before formatting"],"tags":["jax","pallas","mosaic-gpu","sparse","tcgen05","shape","rank"],"backgroundTag":"wrong-array-shape","analyzedSha":"1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb","analyzedAt":"2026-08-27T09:53:25.647Z","schemaVersion":2},"datasetVersion":"2026-08-27T13:17:12.746Z"}