python3Packages.numpyro: 0.20.1 -> 0.21.0

Diff: https://github.com/pyro-ppl/numpyro/compare/0.20.1...0.21.0

Changelog: https://github.com/pyro-ppl/numpyro/releases/tag/0.21.0
This commit is contained in:
Gaetan Lepage
2026-05-05 22:26:32 +00:00
parent 53451ffca0
commit 07bf76aaa3
2 changed files with 3 additions and 27 deletions
@@ -30,21 +30,17 @@
buildPythonPackage (finalAttrs: {
pname = "numpyro";
version = "0.20.1";
version = "0.21.0";
pyproject = true;
__structuredAttrs = true;
src = fetchFromGitHub {
owner = "pyro-ppl";
repo = "numpyro";
tag = finalAttrs.version;
hash = "sha256-sNqllL9nBwXp0kn+HAjvIaHf7LR0UKh9q7DZ20yCr5A=";
hash = "sha256-4NA1m2N0AZy3ausAZc6+PPw175joGC7WwfZr0Ri0uK8=";
};
patches = [
# Remove usage of xla_pmap_p which was removed in jax 0.10.0
./fix-jax-0.10.0-compat.patch
];
build-system = [ setuptools ];
dependencies = [
@@ -1,20 +0,0 @@
diff --git a/numpyro/ops/provenance.py b/numpyro/ops/provenance.py
index 1234567..abcdefg 100644
--- a/numpyro/ops/provenance.py
+++ b/numpyro/ops/provenance.py
@@ -4,7 +4,7 @@
import jax
from jax.api_util import debug_info, flatten_fun, shaped_abstractify
from jax.extend.core import Literal
-from jax.extend.core.primitives import call_p, closed_call_p, jit_p, xla_pmap_p
+from jax.extend.core.primitives import call_p, closed_call_p, jit_p
import jax.extend.linear_util as lu
from jax.interpreters.partial_eval import trace_to_jaxpr_dynamic
@@ -114,7 +114,6 @@ def track_deps_call_rule(eqn, provenance_inputs):
track_deps_rules[call_p] = track_deps_call_rule
-track_deps_rules[xla_pmap_p] = track_deps_call_rule
def track_deps_closed_call_rule(eqn, provenance_inputs):