{"record":{"id":"7c87aa96bde59a99","repo":"jax-ml/jax","slug":"precision-matrix-must-match-data-dims","errorCode":null,"errorMessage":"precision matrix must match data dims","messagePattern":"precision matrix must match data dims","errorType":"exception","errorClass":"ValueError","httpStatus":null,"severity":"error","filePath":"jax/_src/scipy/stats/kde.py","lineNumber":264,"sourceCode":"\n\ndef _gaussian_kernel_convolve(chol, norm, target, weights, mean):\n  diff = target - mean[:, None]\n  alpha = linalg.cho_solve(chol, diff)\n  arg = 0.5 * jnp.sum(diff * alpha, axis=0)\n  return norm * jnp.sum(jnp.exp(-arg) * weights)\n\n\n@api.jit(static_argnums=0)\ndef _gaussian_kernel_eval(in_log, points, values, xi, precision):\n  points, values, xi, precision = promote_dtypes_inexact(\n      points, values, xi, precision)\n  d = points.shape[1]\n\n  if xi.shape[1] != d:\n    raise ValueError(\"points and xi must have same trailing dim\")\n  if precision.shape != (d, d):\n    raise ValueError(\"precision matrix must match data dims\")\n\n  whitening = linalg.cholesky(precision, lower=True)\n  points = jnp.dot(points, whitening)\n  xi = jnp.dot(xi, whitening)\n  log_norm = jnp.sum(jnp.log(\n      jnp.diag(whitening))) - 0.5 * d * jnp.log(2 * np.pi)\n\n  def kernel(x_test, x_train, y_train):\n    arg = log_norm - 0.5 * jnp.sum(jnp.square(x_train - x_test))\n    if in_log:\n      return jnp.log(y_train) + arg\n    else:\n      return y_train * jnp.exp(arg)\n\n  reduce = special.logsumexp if in_log else jnp.sum\n  reduced_kernel = lambda x: reduce(api.vmap(kernel, in_axes=(None, 0, 0))\n                                    (x, points, values),\n                                    axis=0)","sourceCodeStart":246,"sourceCodeEnd":282,"githubUrl":"https://github.com/jax-ml/jax/blob/1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb/jax/_src/scipy/stats/kde.py#L246-L282","documentation":"In _gaussian_kernel_eval the precision matrix must be a full (d, d) matrix matching the data dimensionality, because it is Cholesky-factorized for whitening. Passing a diagonal vector of precisions, a scalar, or a wrongly sized matrix raises this error.","triggerScenarios":"Direct calls to _gaussian_kernel_eval with precision of shape (d,), (1,), or (k, k) with k != d; upstream, this usually surfaces as a wrongly-built inv_cov passed through evaluate.","commonSituations":"Constructing a per-feature precision vector instead of a matrix; using a Cholesky factor or correlation matrix of the wrong size when wiring custom precision into the KDE.","solutions":["Wrap per-dimension precisions into a diagonal matrix: jnp.diag(precision_vec)","Ensure precision.shape == (d, d) where d == points.shape[1]","If precision is a Cholesky factor L, pass L @ L.T"],"exampleFix":"// before\nvals = _gaussian_kernel_eval(False, pts, vals, xi, prec_vec)  # shape (d,)\n// after\nvals = _gaussian_kernel_eval(False, pts, vals, xi, jnp.diag(prec_vec))","handlingStrategy":"validation","validationCode":"d = points.shape[1]\nif precision.ndim == 1:\n    precision = jnp.diag(precision)\nassert precision.shape == (d, d)","typeGuard":"def precision_matches(precision, points) -> bool:\n    d = jnp.asarray(points).shape[1]\n    return jnp.asarray(precision).shape == (d, d)","tryCatchPattern":null,"preventionTips":["Never pass precision vectors; always full (d, d) matrices","Build precision from L @ L.T when starting from a Cholesky factor","Validate matrix shapes in kernel unit tests"],"tags":["jax","scipy","kde","precision-matrix","shape-validation"],"backgroundTag":"invalid-matrix-shape","analyzedSha":"1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb","analyzedAt":"2026-08-27T09:53:25.647Z","schemaVersion":2},"datasetVersion":"2026-08-27T13:17:12.746Z"}