如何在Streamlit中追踪PyMC采样进程并显示进度?
问题
我开发了一个基于Streamlit的项目,可通过浏览器组件设计PyMC模型,运行状态良好。但执行采样器时,从模型编译到采样结束存在较长延迟,希望为用户提供进程正常运行的提示,若能显示采样速度则更佳。我尝试通过回调将PyMC与Streamlit进度条关联,但无法正常工作,恳请提供可行的修复方案或其他追踪进度的实现思路。
附尝试的代码:
import streamlit as st import pymc as pm from pymc.progress_bar import ProgressBarManager st.title("PyMC + Nutpie Sampler") n_draws = 1000 n_tune = 1000 n_chains = 4 if st.button("Run Sampling"): total_steps = n_draws * n_chains chain_draws = {i: 0 for i in range(n_chains)} progress_bar = st.progress(0, text=f"Sampling for {n_chains * n_draws} steps...") old_update = ProgressBarManager.update def new_update(self, chain_idx, is_last, draw, tuning, stats): if not tuning: chain_draws[chain_idx] += 1 completed = sum(chain_draws.values()) progress = min(completed / total_steps, 1.0) progress_bar.progress( progress, text=f"Sampling... {completed}/{total_steps} steps ({progress * 100:.1f}%)" ) old_update(self, chain_idx, is_last, draw, tuning, stats) ProgressBarManager.update = new_update with pm.Model() as model: mu = pm.Normal("mu", mu=0, sigma=1) obs = pm.Normal("obs", mu=mu, sigma=1, observed=[1, 2, 3]) trace = pm.sample( draws=n_draws, tune=n_tune, chains=n_chains, nuts_sampler="nutpie", nuts_sampler_kwargs={"backend": "jax", "gradient_backend": "jax"}, ) # Restore original to avoid side effects on reruns ProgressBarManager.update = old_update progress_bar.progress(1.0, text="Sampling complete!") st.success("Sampling complete!") st.subheader("Posterior Summary") st.dataframe(pm.stats.summary(trace))
原因分析
问题核心在于三点:
- Nutpie采样器不兼容PyMC默认进度管理:Nutpie作为第三方采样器,未使用PyMC内置的
ProgressBarManager逻辑,直接替换其update方法无法触发回调。 - Streamlit UI更新的线程限制:Streamlit组件必须在主线程更新,而PyMC采样多在子线程/进程执行,回调内直接更新进度条会因线程隔离失效。
- 未覆盖全流程进度:原代码仅统计采样阶段(非调参)的步数,编译、调参阶段的延迟无提示,用户感知不到进程状态。
修复方案:Nutpie原生回调+线程安全UI更新
Nutpie支持通过callback参数传入自定义回调函数,结合Streamlit的容器组件和时间统计,可实现全流程进度追踪及采样速度显示:
import streamlit as st import pymc as pm import time from collections import defaultdict st.title("PyMC + Nutpie Sampler") n_draws = 1000 n_tune = 1000 n_chains = 4 if st.button("Run Sampling"): # 初始化进度追踪状态 total_sample_steps = n_draws * n_chains total_tune_steps = n_tune * n_chains chain_progress = defaultdict(lambda: {"tune": 0, "sample": 0}) start_time = time.time() last_update_time = start_time last_completed_samples = 0 # 创建灵活的进度容器 progress_container = st.empty() def nutpie_callback(chain_idx, stage, **kwargs): nonlocal last_update_time, last_completed_samples current_time = time.time() # 更新对应链的阶段进度 if stage == "tune": chain_progress[chain_idx]["tune"] += 1 elif stage == "sample": chain_progress[chain_idx]["sample"] += 1 # 计算已完成总步数 completed_tune = sum(p["tune"] for p in chain_progress.values()) completed_sample = sum(p["sample"] for p in chain_progress.values()) # 计算采样速度(每0.5秒更新一次避免频繁刷新) speed = 0 if completed_sample > 0 and current_time - last_update_time > 0.5: speed = (completed_sample - last_completed_samples) / (current_time - last_update_time) last_update_time = current_time last_completed_samples = completed_sample # 更新UI(确保在主线程上下文执行) with progress_container: if completed_tune < total_tune_steps: # 调参阶段提示 progress = completed_tune / total_tune_steps st.progress(progress, text=f"调参中... {completed_tune}/{total_tune_steps} 步 ({progress*100:.1f}%)") else: # 采样阶段提示(含速度) progress = completed_sample / total_sample_steps speed_text = f" | 速度: {speed:.1f} 步/秒" if speed > 0 else "" st.progress(progress, text=f"采样中... {completed_sample}/{total_sample_steps} 步 ({progress*100:.1f}%){speed_text}") # 模型定义与采样 with st.spinner("模型编译中..."): with pm.Model() as model: mu = pm.Normal("mu", mu=0, sigma=1) obs = pm.Normal("obs", mu=mu, sigma=1, observed=[1, 2, 3]) trace = pm.sample( draws=n_draws, tune=n_tune, chains=n_chains, nuts_sampler="nutpie", nuts_sampler_kwargs={ "backend": "jax", "gradient_backend": "jax", "callback": nutpie_callback }, ) # 采样完成状态更新 with progress_container: st.progress(1.0, text="采样完成!") st.success("采样完成!") st.subheader("后验分布汇总") st.dataframe(pm.stats.summary(trace))
其他实现思路
- 后台线程分离任务:将采样逻辑放到后台线程,主线程定期读取进度状态并更新UI,避免阻塞Streamlit渲染循环。
- PyMC默认采样器适配:若使用PyMC原生NUTS采样器,可结合
pm.callbacks.CheckParametersConvergence扩展自定义进度追踪逻辑。 - 阶段式加载提示:在编译、调参、采样各阶段分别显示针对性提示,覆盖全流程等待时间。
内容的提问来源于stack exchange,提问作者Vital Fernández
相关产品推荐
相关产品推荐

