From 1cff290e60574956951af05e205d6a5a328e51c6 Mon Sep 17 00:00:00 2001 From: Karthik Laishetti Date: Tue, 21 Jul 2026 11:58:53 +0530 Subject: [PATCH 01/10] Update progress_bar handling in jax.py for blackjax Handle progress_bar parameter for blackjax version compatibility. --- pymc/sampling/jax.py | 24 +++++++++++++++++++----- 1 file changed, 19 insertions(+), 5 deletions(-) diff --git a/pymc/sampling/jax.py b/pymc/sampling/jax.py index 4305d3bdf6..2e3db9a7f5 100644 --- a/pymc/sampling/jax.py +++ b/pymc/sampling/jax.py @@ -250,6 +250,12 @@ def _blackjax_inference_loop( else: raise ValueError("Only supporting 'nuts' or 'hmc' as algorithm to draw samples.") + # Must be popped before calling window_adaptation: blackjax >= 1.6 removed + # the progress_bar parameter from window_adaptation and forwards any + # unrecognized kwargs straight into the NUTS kernel, which raises + # TypeError. See https://github.com/pymc-devs/pymc/issues/8367. + progress_bar = adaptation_kwargs.pop("progress_bar", False) + adapt = blackjax.window_adaptation( algorithm=algorithm, logdensity_fn=logp_fn, @@ -274,12 +280,20 @@ def _one_step(state, xs): } return state, (position, stats) - progress_bar = adaptation_kwargs.pop("progress_bar", False) - keys = jax.random.split(seed, draws) - scan_fn = blackjax.progress_bar.gen_scan_fn(draws, progress_bar) - _, (samples, stats) = scan_fn(_one_step, last_state, (jnp.arange(draws), keys)) - + if hasattr(blackjax.progress_bar, "gen_scan_fn"): + # blackjax < 1.6: progress_bar is a module exposing gen_scan_fn, + # which wraps jax.lax.scan directly. + scan_fn = blackjax.progress_bar.gen_scan_fn(draws, progress_bar) + _, (samples, stats) = scan_fn(_one_step, last_state, (jnp.arange(draws), keys)) + elif progress_bar: + # blackjax >= 1.6: progress_bar is a context manager that + # monkeypatches jax.lax.scan for its duration instead. + with blackjax.progress_bar(label="NUTS"): + _, (samples, stats) = jax.lax.scan(_one_step, last_state, (jnp.arange(draws), keys)) + else: + _, (samples, stats) = jax.lax.scan(_one_step, last_state, (jnp.arange(draws), keys)) + return samples, stats From 49423d558add21544390d3eb5ee0d8808f96898d Mon Sep 17 00:00:00 2001 From: Karthik Laishetti Date: Tue, 21 Jul 2026 12:02:28 +0530 Subject: [PATCH 02/10] Implement tests for blackjax NUTS progress bar handling Add regression tests for blackjax progress bar compatibility. --- tests/sampling/test_jax.py | 75 ++++++++++++++++++++++++++++++++++++++ 1 file changed, 75 insertions(+) diff --git a/tests/sampling/test_jax.py b/tests/sampling/test_jax.py index 3006fd2040..ebb2589d5f 100644 --- a/tests/sampling/test_jax.py +++ b/tests/sampling/test_jax.py @@ -502,3 +502,78 @@ def test_convergence_warnings(caplog, nuts_sampler): [record] = caplog.records assert re.match(r"There were \d+ divergences after tuning", record.message) + +class TestBlackjaxProgressBarCompat: + """Regression tests for https://github.com/pymc-devs/pymc/issues/8367. + + blackjax >= 1.6 removed the progress_bar parameter from + window_adaptation (unknown kwargs get forwarded straight into the NUTS + kernel and raise TypeError) and replaced the progress_bar module + (with its gen_scan_fn helper) with a context-manager function. + """ + + def test_sample_blackjax_nuts_progressbar_false(self): + # This is the exact reproduction from the issue: progress_bar + # reaching window_adaptation used to raise + # "TypeError: ... got an unexpected keyword argument 'progress_bar'" + # on any blackjax >= 1.6. + with pm.Model(): + x = pm.Normal("x", 0.0, 1.0) + pm.Normal("obs", x, 1.0, observed=np.array([0.3, -0.1, 0.5])) + idata = pm.sample( + draws=10, + tune=10, + chains=1, + cores=1, + nuts_sampler="blackjax", + progressbar=False, + ) + assert idata.posterior["x"].shape == (1, 10) + + def test_sample_blackjax_nuts_progressbar_true(self): + # Exercises the progress_bar=True path specifically, which on + # blackjax >= 1.6 (after fixing the TypeError above) used to hit a + # second break: AttributeError, since blackjax.progress_bar is no + # longer a module with a gen_scan_fn attribute. + with pm.Model(): + x = pm.Normal("x", 0.0, 1.0) + pm.Normal("obs", x, 1.0, observed=np.array([0.3, -0.1, 0.5])) + idata = pm.sample( + draws=10, + tune=10, + chains=1, + cores=1, + nuts_sampler="blackjax", + progressbar=True, + ) + assert idata.posterior["x"].shape == (1, 10) + + def test_progress_bar_popped_before_window_adaptation(self): + # Directly asserts the ordering fix: progress_bar must never reach + # blackjax.window_adaptation's kwargs, regardless of blackjax + # version/API shape. + import blackjax + + from pymc.sampling.jax import _blackjax_inference_loop + + original_window_adaptation = blackjax.window_adaptation + + def spy_window_adaptation(*args, **kwargs): + assert "progress_bar" not in kwargs, "progress_bar leaked into window_adaptation kwargs" + return original_window_adaptation(*args, **kwargs) + + with pm.Model() as model: + x = pm.Normal("x", 0.0, 1.0) + pm.Normal("obs", x, 1.0, observed=np.array([0.3, -0.1, 0.5])) + logp_fn = get_jaxified_logp(model) + + with mock.patch("blackjax.window_adaptation", side_effect=spy_window_adaptation): + _blackjax_inference_loop( + seed=jax.random.PRNGKey(0), + init_position=[np.array(0.0)], + logp_fn=logp_fn, + draws=5, + tune=5, + target_accept=0.8, + progress_bar=False, + ) From 088c2419938cec889b2b494f2a49afba75a089ce Mon Sep 17 00:00:00 2001 From: Karthik Laishetti Date: Tue, 21 Jul 2026 12:13:47 +0530 Subject: [PATCH 03/10] Remove unnecessary blank line in jax.py --- pymc/sampling/jax.py | 1 - 1 file changed, 1 deletion(-) diff --git a/pymc/sampling/jax.py b/pymc/sampling/jax.py index 2e3db9a7f5..2d264f232f 100644 --- a/pymc/sampling/jax.py +++ b/pymc/sampling/jax.py @@ -293,7 +293,6 @@ def _one_step(state, xs): _, (samples, stats) = jax.lax.scan(_one_step, last_state, (jnp.arange(draws), keys)) else: _, (samples, stats) = jax.lax.scan(_one_step, last_state, (jnp.arange(draws), keys)) - return samples, stats From 5912a4fc66c825512784e77428b9b0817e1f6553 Mon Sep 17 00:00:00 2001 From: Karthik Laishetti Date: Wed, 22 Jul 2026 10:38:04 +0530 Subject: [PATCH 04/10] Add progress bar support for blackjax sampling --- pymc/sampling/jax.py | 15 +++++++++++++++ 1 file changed, 15 insertions(+) diff --git a/pymc/sampling/jax.py b/pymc/sampling/jax.py index 2d264f232f..b37fe9ea33 100644 --- a/pymc/sampling/jax.py +++ b/pymc/sampling/jax.py @@ -295,6 +295,21 @@ def _one_step(state, xs): _, (samples, stats) = jax.lax.scan(_one_step, last_state, (jnp.arange(draws), keys)) return samples, stats + keys = jax.random.split(seed, draws) + if hasattr(blackjax.progress_bar, "gen_scan_fn"): + # blackjax < 1.6: progress_bar is a module exposing gen_scan_fn, + # which wraps jax.lax.scan directly. + scan_fn = blackjax.progress_bar.gen_scan_fn(draws, progress_bar) + _, (samples, stats) = scan_fn(_one_step, last_state, (jnp.arange(draws), keys)) + elif progress_bar: + # blackjax >= 1.6: progress_bar is a context manager that + # monkeypatches jax.lax.scan for its duration instead. + with blackjax.progress_bar(label="NUTS"): + _, (samples, stats) = jax.lax.scan(_one_step, last_state, (jnp.arange(draws), keys)) + else: + _, (samples, stats) = jax.lax.scan(_one_step, last_state, (jnp.arange(draws), keys)) + return samples, stats + def _sample_blackjax_nuts( model: Model, From 63b53b6b74bc1ed0ce0e71c4053d20e61485d8fa Mon Sep 17 00:00:00 2001 From: "pre-commit-ci[bot]" <66853113+pre-commit-ci[bot]@users.noreply.github.com> Date: Wed, 22 Jul 2026 05:10:05 +0000 Subject: [PATCH 05/10] [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci --- tests/sampling/test_jax.py | 1 + 1 file changed, 1 insertion(+) diff --git a/tests/sampling/test_jax.py b/tests/sampling/test_jax.py index ebb2589d5f..aba2c4ca6f 100644 --- a/tests/sampling/test_jax.py +++ b/tests/sampling/test_jax.py @@ -503,6 +503,7 @@ def test_convergence_warnings(caplog, nuts_sampler): [record] = caplog.records assert re.match(r"There were \d+ divergences after tuning", record.message) + class TestBlackjaxProgressBarCompat: """Regression tests for https://github.com/pymc-devs/pymc/issues/8367. From 30c25e0efd9477400f5ed9f612d2a212a7cea5a9 Mon Sep 17 00:00:00 2001 From: Karthik Laishetti Date: Sat, 8 Aug 2026 13:51:46 +0530 Subject: [PATCH 06/10] Improve progress bar integration in jax.py Refactor progress bar handling for blackjax sampling. --- pymc/sampling/jax.py | 30 +++++++++++------------------- 1 file changed, 11 insertions(+), 19 deletions(-) diff --git a/pymc/sampling/jax.py b/pymc/sampling/jax.py index b37fe9ea33..562ce6c208 100644 --- a/pymc/sampling/jax.py +++ b/pymc/sampling/jax.py @@ -281,36 +281,28 @@ def _one_step(state, xs): return state, (position, stats) keys = jax.random.split(seed, draws) - if hasattr(blackjax.progress_bar, "gen_scan_fn"): + if hasattr(blackjax.progress_bar, "gen_scan_fn"): # blackjax < 1.6: progress_bar is a module exposing gen_scan_fn, # which wraps jax.lax.scan directly. scan_fn = blackjax.progress_bar.gen_scan_fn(draws, progress_bar) _, (samples, stats) = scan_fn(_one_step, last_state, (jnp.arange(draws), keys)) elif progress_bar: - # blackjax >= 1.6: progress_bar is a context manager that - # monkeypatches jax.lax.scan for its duration instead. - with blackjax.progress_bar(label="NUTS"): + try: + with blackjax.progress_bar(label="NUTS"): + _, (samples, stats) = jax.lax.scan(_one_step, last_state, (jnp.arange(draws), keys)) + except ImportError: + warnings.warn( + "blackjax progress bar needs the 'progress' extra " + "(pip install 'blackjax[progress]'); sampling without it.", + UserWarning, + stacklevel=2, + ) _, (samples, stats) = jax.lax.scan(_one_step, last_state, (jnp.arange(draws), keys)) else: _, (samples, stats) = jax.lax.scan(_one_step, last_state, (jnp.arange(draws), keys)) - return samples, stats - keys = jax.random.split(seed, draws) - if hasattr(blackjax.progress_bar, "gen_scan_fn"): - # blackjax < 1.6: progress_bar is a module exposing gen_scan_fn, - # which wraps jax.lax.scan directly. - scan_fn = blackjax.progress_bar.gen_scan_fn(draws, progress_bar) - _, (samples, stats) = scan_fn(_one_step, last_state, (jnp.arange(draws), keys)) - elif progress_bar: - # blackjax >= 1.6: progress_bar is a context manager that - # monkeypatches jax.lax.scan for its duration instead. - with blackjax.progress_bar(label="NUTS"): - _, (samples, stats) = jax.lax.scan(_one_step, last_state, (jnp.arange(draws), keys)) - else: - _, (samples, stats) = jax.lax.scan(_one_step, last_state, (jnp.arange(draws), keys)) return samples, stats - def _sample_blackjax_nuts( model: Model, target_accept: float, From 34ddf8f31b9965aa78667b5973e23a409635946d Mon Sep 17 00:00:00 2001 From: Karthik Laishetti Date: Sat, 8 Aug 2026 08:32:15 +0000 Subject: [PATCH 07/10] fix: correct indentation on line 284 --- pymc/sampling/jax.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/pymc/sampling/jax.py b/pymc/sampling/jax.py index 562ce6c208..bd3994ab13 100644 --- a/pymc/sampling/jax.py +++ b/pymc/sampling/jax.py @@ -281,7 +281,7 @@ def _one_step(state, xs): return state, (position, stats) keys = jax.random.split(seed, draws) - if hasattr(blackjax.progress_bar, "gen_scan_fn"): + if hasattr(blackjax.progress_bar, "gen_scan_fn"): # blackjax < 1.6: progress_bar is a module exposing gen_scan_fn, # which wraps jax.lax.scan directly. scan_fn = blackjax.progress_bar.gen_scan_fn(draws, progress_bar) From 0f550d9ef1b005a5e6da95a181facb8bb743ebf9 Mon Sep 17 00:00:00 2001 From: Karthik Laishetti Date: Sat, 8 Aug 2026 08:33:20 +0000 Subject: [PATCH 08/10] style: apply ruff format --- pymc/sampling/jax.py | 1 + 1 file changed, 1 insertion(+) diff --git a/pymc/sampling/jax.py b/pymc/sampling/jax.py index bd3994ab13..67b1d19cea 100644 --- a/pymc/sampling/jax.py +++ b/pymc/sampling/jax.py @@ -303,6 +303,7 @@ def _one_step(state, xs): return samples, stats + def _sample_blackjax_nuts( model: Model, target_accept: float, From 0a13398ed19ba4dd9e1c0caa2de1b9216d88f986 Mon Sep 17 00:00:00 2001 From: Karthik Laishetti Date: Sun, 9 Aug 2026 13:37:33 +0530 Subject: [PATCH 09/10] Fix progress bar leakage in window adaptation test Refactor test for progress bar handling in JAX sampling. --- tests/sampling/test_jax.py | 9 +++------ 1 file changed, 3 insertions(+), 6 deletions(-) diff --git a/tests/sampling/test_jax.py b/tests/sampling/test_jax.py index aba2c4ca6f..107de62667 100644 --- a/tests/sampling/test_jax.py +++ b/tests/sampling/test_jax.py @@ -19,6 +19,7 @@ from typing import Any from unittest import mock +import blackjax import jax import numpy as np import pytensor @@ -31,8 +32,8 @@ import pymc as pm -from pymc.exceptions import ImputationWarning from pymc.sampling.jax import ( + _blackjax_inference_loop, _get_batched_jittered_initial_points, _get_log_likelihood, _replace_shared_variables, @@ -549,14 +550,10 @@ def test_sample_blackjax_nuts_progressbar_true(self): ) assert idata.posterior["x"].shape == (1, 10) - def test_progress_bar_popped_before_window_adaptation(self): + def test_progress_bar_popped_before_window_adaptation(self): # Directly asserts the ordering fix: progress_bar must never reach # blackjax.window_adaptation's kwargs, regardless of blackjax # version/API shape. - import blackjax - - from pymc.sampling.jax import _blackjax_inference_loop - original_window_adaptation = blackjax.window_adaptation def spy_window_adaptation(*args, **kwargs): From 28a0ed06b6baefa2f5b698d87baecc1fea3490c9 Mon Sep 17 00:00:00 2001 From: Karthik Laishetti Date: Mon, 10 Aug 2026 02:48:09 +0000 Subject: [PATCH 10/10] fix: remove duplicate pytensor import, restore missing ImputationWarning import, add trailing newline --- tests/sampling/test_jax.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/tests/sampling/test_jax.py b/tests/sampling/test_jax.py index 107de62667..e6195c824f 100644 --- a/tests/sampling/test_jax.py +++ b/tests/sampling/test_jax.py @@ -32,6 +32,7 @@ import pymc as pm +from pymc.exceptions import ImputationWarning from pymc.sampling.jax import ( _blackjax_inference_loop, _get_batched_jittered_initial_points, @@ -550,7 +551,7 @@ def test_sample_blackjax_nuts_progressbar_true(self): ) assert idata.posterior["x"].shape == (1, 10) - def test_progress_bar_popped_before_window_adaptation(self): + def test_progress_bar_popped_before_window_adaptation(self): # Directly asserts the ordering fix: progress_bar must never reach # blackjax.window_adaptation's kwargs, regardless of blackjax # version/API shape.