From 07bf76aaa3a962d3ae92a936cfdc94215ce0340f Mon Sep 17 00:00:00 2001 From: Gaetan Lepage Date: Mon, 4 May 2026 22:47:26 +0000 Subject: [PATCH] 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 --- .../python-modules/numpyro/default.nix | 10 +++------- .../numpyro/fix-jax-0.10.0-compat.patch | 20 ------------------- 2 files changed, 3 insertions(+), 27 deletions(-) delete mode 100644 pkgs/development/python-modules/numpyro/fix-jax-0.10.0-compat.patch diff --git a/pkgs/development/python-modules/numpyro/default.nix b/pkgs/development/python-modules/numpyro/default.nix index 2e114b7b5e06..55e737354728 100644 --- a/pkgs/development/python-modules/numpyro/default.nix +++ b/pkgs/development/python-modules/numpyro/default.nix @@ -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 = [ diff --git a/pkgs/development/python-modules/numpyro/fix-jax-0.10.0-compat.patch b/pkgs/development/python-modules/numpyro/fix-jax-0.10.0-compat.patch deleted file mode 100644 index ad9e937a1e06..000000000000 --- a/pkgs/development/python-modules/numpyro/fix-jax-0.10.0-compat.patch +++ /dev/null @@ -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):