You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何在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))
原因分析

问题核心在于三点:

  1. Nutpie采样器不兼容PyMC默认进度管理:Nutpie作为第三方采样器,未使用PyMC内置的ProgressBarManager逻辑,直接替换其update方法无法触发回调。
  2. Streamlit UI更新的线程限制:Streamlit组件必须在主线程更新,而PyMC采样多在子线程/进程执行,回调内直接更新进度条会因线程隔离失效。
  3. 未覆盖全流程进度:原代码仅统计采样阶段(非调参)的步数,编译、调参阶段的延迟无提示,用户感知不到进程状态。
修复方案: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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.01 23:14:52