{"record":{"id":"a18a297235e06f5b","repo":"jax-ml/jax","slug":"unsupported-ndim-x-ndim","errorCode":null,"errorMessage":"Unsupported ndim: {x.ndim}","messagePattern":"Unsupported ndim: (.+?)","errorType":"validation","errorClass":"ValueError","httpStatus":null,"severity":"error","filePath":"jax/_src/pallas/mosaic_gpu/core.py","lineNumber":1285,"sourceCode":"    return pp.text(f\"{{unswizzle({self.swizzle})}}\")\n\n\n@tree_util.register_dataclass\n@dataclasses.dataclass(frozen=True)\nclass CollapseLeadingBatchDimensionsTransform(state_types.Transform):\n  \"\"\"A transform that collapses leading batch dimensions into the minor dimension.\n\n  Specifically, it maps `(*batch_shape, m, n)` to `(m, math.prod(batch_shape) *\n  n)`.\n  \"\"\"\n\n  def transform_type(\n      self, x: jax_core.AbstractValue\n  ) -> state_types.AbstractRef:\n    match x:\n      case jax_core.ShapedArray():\n        if x.ndim < 2:\n          raise ValueError(f\"Unsupported ndim: {x.ndim}\")\n        batch_size = math.prod(x.shape[:-2])\n        transformed_shape = (x.shape[-2], batch_size * x.shape[-1])\n        return x.update(shape=transformed_shape)\n      case state_types.AbstractRef():\n        return x.update(inner_aval=self.transform_type(x.inner_aval))\n      case _:\n        raise TypeError(f\"Unsupported type: {x}\")\n\n  def undo(self, x: jax_core.AbstractValue) -> state_types.Transform:\n    assert hasattr(x, \"shape\")\n    return ExpandLeadingBatchDimensionsTransform(x.shape[:-2])\n\n\n@tree_util.register_dataclass\n@dataclasses.dataclass(frozen=True)\nclass ExpandLeadingBatchDimensionsTransform(state_types.Transform):\n  \"\"\"The inverse of CollapseLeadingBatchDimensionsTransform.\n","sourceCodeStart":1267,"sourceCodeEnd":1303,"githubUrl":"https://github.com/jax-ml/jax/blob/1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb/jax/_src/pallas/mosaic_gpu/core.py#L1267-L1303","documentation":"ExpandLeadingBatchDimensionsTransform.transform_type requires the input ShapedArray to have ndim >= 2, because it folds all leading dims into a batch factor of a 2-D (rows x cols) physical layout. Arrays with 0 or 1 dims cannot be represented and raise ValueError('Unsupported ndim').","triggerScenarios":"Passing a scalar (ndim 0) or 1-D array through a transform that expands leading batch dimensions, e.g. mapping a 1-D vector ref to a 2-D swizzled buffer via get_ref_aval/to_block_mapping.","commonSituations":"Trying to store 1-D bias/scale vectors in a 2-D-only WGMMA layout inside Mosaic kernels; allocating scalars in transforms meant for matrices.","solutions":["Reshape to at least 2-D before applying the transform (e.g. x[None, :])","Store 1-D data in a separate non-transformed buffer","Pad the vector to a (1, n) matrix"],"exampleFix":"// before\nref = make_ref(vec)  # vec.ndim == 1 -> ValueError\n// after\nref = make_ref(vec[None, :])  # shape (1, n)","handlingStrategy":"validation","validationCode":"assert x.ndim >= 2, f'need ndim >= 2, got {x.ndim}'","typeGuard":"def is_batch_expandable(x) -> bool:\\n    return getattr(x, 'ndim', 0) >= 2","tryCatchPattern":null,"preventionTips":["Reshape 1-D/scalar data to 2-D before batch-expanding transforms"],"tags":["jax","pallas","ndim","shape-validation"],"backgroundTag":"invalid-shape-for-transform","analyzedSha":"1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb","analyzedAt":"2026-08-27T09:53:25.647Z","schemaVersion":2},"datasetVersion":"2026-08-27T13:17:12.746Z"}