{"record":{"id":"04778d747d2cca7a","repo":"jax-ml/jax","slug":"m-has-more-than-2-dimensions","errorCode":null,"errorMessage":"m has more than 2 dimensions","messagePattern":"m has more than 2 dimensions","errorType":"exception","errorClass":"ValueError","httpStatus":null,"severity":"error","filePath":"jax/_src/numpy/lax_numpy.py","lineNumber":9185,"sourceCode":"    points drawn from a 3-dimensional standard normal distribution:\n\n    >>> key = jax.random.key(0)\n    >>> x = jax.random.normal(key, shape=(3, 100))\n    >>> with jnp.printoptions(precision=2):\n    ...   print(jnp.cov(x))\n    [[0.9  0.03 0.1 ]\n     [0.03 1.   0.01]\n     [0.1  0.01 0.85]]\n  \"\"\"\n  if y is not None:\n    m, y = util.promote_args_inexact(\"cov\", m, y)\n    if y.ndim > 2:\n      raise ValueError(\"y has more than 2 dimensions\")\n  else:\n    m, = util.promote_args_inexact(\"cov\", m)\n\n  if m.ndim > 2:\n    raise ValueError(\"m has more than 2 dimensions\")  # same as numpy error\n\n  if dtype is not None and not dtypes.issubdtype(dtype, np.inexact):\n    raise ValueError(f\"cov: dtype must be a subclass of float or complex; got {dtype=}\")\n\n  X = atleast_2d(m)\n  if not rowvar and m.ndim != 1:\n    X = X.T\n  if X.shape[0] == 0:\n    return array([]).reshape(0, 0)\n\n  if y is not None:\n    y_arr = atleast_2d(y)\n    if not rowvar and y_arr.shape[0] != 1:\n      y_arr = y_arr.T\n    X = concatenate((X, y_arr), axis=0)\n  if X.shape[1] == 0:\n    cov_shape = () if X.shape[0] == 1 else (X.shape[0], X.shape[0])\n    return array_creation.full(cov_shape, np.nan, dtype=X.dtype)","sourceCodeStart":9167,"sourceCodeEnd":9203,"githubUrl":"https://github.com/jax-ml/jax/blob/1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb/jax/_src/numpy/lax_numpy.py#L9167-L9203","documentation":"jnp.cov requires the primary data matrix m to be 1D (a single variable's observations) or 2D (variables x observations). If m.ndim > 2 after promotion, ValueError('m has more than 2 dimensions') is raised, matching NumPy's behavior.","triggerScenarios":"Calling jnp.cov on a 3D tensor such as shape (batch, features, time), e.g. jnp.cov(images) where images.ndim == 3.","commonSituations":"Applying cov to batches of images, video, or windowed time-series without flattening; migrating NumPy pipelines that already reshaped data but losing the reshape step.","solutions":["Flatten leading dimensions: m.reshape(-1, m.shape[-1]) or m.reshape(m.shape[0], -1) depending on variable layout","Use jax.vmap(jnp.cov) to get per-sample covariance matrices for batched data","Restructure data so rows are variables and columns observations"],"exampleFix":"// before\nc = jnp.cov(batch_3d)  # ValueError\n// after\nc = jax.vmap(jnp.cov)(batch_3d)  # or reshape to 2D","handlingStrategy":"validation","validationCode":"m = jnp.asarray(m)\nassert m.ndim <= 2, 'm must be 1D or 2D for cov'\njnp.cov(m)","typeGuard":"def is_cov_input(x) -> bool:\n    return 1 <= jnp.asarray(x).ndim <= 2","tryCatchPattern":null,"preventionTips":["Flatten batch dims with reshape","vmap(jnp.cov) over batched tensors","Remember rows=variables, cols=observations"],"tags":["jax","statistics","covariance","ndim-validation"],"backgroundTag":"shape-validation-failed","analyzedSha":"1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb","analyzedAt":"2026-08-27T09:53:25.647Z","schemaVersion":2},"datasetVersion":"2026-08-27T13:17:12.746Z"}