{"record":{"id":"ee3cd13c08009f04","repo":"jax-ml/jax","slug":"the-out-argument-to-jnp-nanvar-is-not-supported","errorCode":null,"errorMessage":"The 'out' argument to jnp.nanvar is not supported.","messagePattern":"The 'out' argument to jnp\\.nanvar is not supported\\.","errorType":"validation","errorClass":"NotImplementedError","httpStatus":null,"severity":"error","filePath":"jax/_src/numpy/reductions.py","lineNumber":1919,"sourceCode":"     [ 4.  ]]\n\n    To include specific elements of the array to compute variance, you can use\n    ``where``.\n\n    >>> where = jnp.array([[1, 0, 1, 0],\n    ...                    [0, 1, 1, 0],\n    ...                    [1, 1, 0, 1]], dtype=bool)\n    >>> jnp.nanvar(x, axis=1, keepdims=True, where=where)\n    Array([[2.25],\n           [0.  ],\n           [4.  ]], dtype=float32)\n  \"\"\"\n  a = ensure_arraylike(\"nanvar\", a)\n  where = check_where(\"nanvar\", where)\n  if dtype is not None:\n    dtype = dtypes.check_and_canonicalize_user_dtype(dtype, \"nanvar\")\n  if out is not None:\n    raise NotImplementedError(\"The 'out' argument to jnp.nanvar is not supported.\")\n  return _nanvar(a, axis=axis, dtype=dtype, out=out, ddof=ddof, keepdims=keepdims, where=where, a_mean=mean)\n\ndef _nanvar(a: Array, *, axis: Axis = None, dtype: DTypeLike | None = None, out: None = None,\n           ddof: int = 0, keepdims: bool = False,\n           where: ArrayLike | None = None, a_mean: ArrayLike | None = None) -> Array:\n  computation_dtype, dtype = _var_promote_types(a.dtype, dtype)\n  a = lax.asarray(a).astype(computation_dtype)\n  if a_mean is None:\n    a_mean = nanmean(a, axis, dtype=computation_dtype, keepdims=True, where=where)\n  else:\n    a_mean = ensure_arraylike(\"nanvar\", a_mean).astype(computation_dtype)\n\n  centered = _where(lax._isnan(a), 0, lax.sub(a, a_mean))  # double-where trick for gradients.\n  if dtypes.issubdtype(centered.dtype, np.complexfloating):\n    centered = lax.real(lax.mul(centered, lax.conj(centered)))\n  else:\n    centered = lax.square(centered)\n","sourceCodeStart":1901,"sourceCodeEnd":1937,"githubUrl":"https://github.com/jax-ml/jax/blob/1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb/jax/_src/numpy/reductions.py#L1901-L1937","documentation":"jnp.nanvar rejects the out= argument (NotImplementedError) since JAX cannot write results into a caller-provided mutable buffer. The parameter is retained purely for NumPy API compatibility.","triggerScenarios":"Calling jnp.nanvar(a, out=buf); also reached indirectly if nanstd forwards out-related state — nanstd calls nanvar internally.","commonSituations":"Ported NumPy variance-with-NaNs code that used out=; kwargs passthrough from a config-driven stats pipeline.","solutions":["Remove out= and use the returned array","If allocation churn matters, wrap the call in jax.jit so the compiler handles buffer reuse"],"exampleFix":"// before\nnp.nanvar(x, out=buf)\n// after\nbuf = jnp.nanvar(x)","handlingStrategy":"validation","validationCode":"assert out is None\nv = jnp.nanvar(a, axis=axis, ddof=ddof, where=where)","typeGuard":null,"tryCatchPattern":null,"preventionTips":["Drop out= from all ported variance calls","Wrap jnp reductions in a small adapter that whitelists kwargs"],"tags":["jax","numpy","nanvar","out-argument"],"backgroundTag":"unsupported-out-argument","analyzedSha":"1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb","analyzedAt":"2026-08-27T09:53:25.647Z","schemaVersion":2},"datasetVersion":"2026-08-27T13:17:12.746Z"}