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

如何在Streamlit中缓存SHAP绘图以解决重复生成耗时问题

问题原因
  • 你使用的0.84.1版本Streamlit的st.cache装饰器会对函数依赖的所有全局变量、输入输出对象执行哈希校验,你函数中用到的shapvs、prep_train都是全量训练集级别的大对象,哈希本身耗时极长,看起来就像函数无限运行。
  • matplotlib的Figure对象属于不可哈希的复杂可变对象,st.cache对这类对象的哈希校验逻辑容易失效,导致缓存命中判定失败,反复重跑函数。
  • shap的summary_plot默认会直接修改当前激活的matplotlib画布状态,这类状态变更不会被st.cache识别,进一步加剧校验失败问题。
可行的缓存方案

推荐优先用缓存图片二进制流的方案,完全规避复杂对象的哈希问题,性能最优:

import io
import matplotlib.pyplot as plt

@st.cache(allow_output_mutation=True, suppress_st_warning=True)
def get_cached_summary_plot():
    fig, ax = plt.subplots(figsize=(12, 18))
    # 必须加show=False,避免shap直接输出画布导致缓存逻辑异常
    shap.summary_plot(shapvs[1], prep_train.iloc[:, :-1].values, 
                     prep_train.columns, max_display=50, show=False)
    # 将画布写入内存缓冲区
    buf = io.BytesIO()
    plt.savefig(buf, format='png', bbox_inches='tight')
    buf.seek(0)
    # 主动关闭画布避免内存泄漏
    plt.close(fig)
    return buf

# 直接读取缓存的二进制流展示图片
st.image(get_cached_summary_plot())

如果确实需要保留Figure对象做后续修改,也可以通过指定哈希规则跳过大数据集的哈希校验,稳定性略低于第一种方案,代码参考:

import pandas as pd

@st.cache(allow_output_mutation=True, 
          hash_funcs={type(shapvs): lambda x: id(x),
                      pd.DataFrame: lambda x: id(x)})
def summary_plot_all():
    fig, axes = plt.subplots(nrows=1, ncols=1, figsize=(12,18))
    shap.summary_plot(shapvs[1], prep_train.iloc[:, :-1].values, 
    prep_train.columns, max_display=50, show=False)
    return fig
    
st.pyplot(summary_plot_all())
plt.close()

该方案通过id()直接取全局变量的内存地址作为哈希值,只要全局变量没有被重新赋值,就会命中缓存,避免了大数据集的哈希耗时。

内容的提问来源于stack exchange,提问作者vpvinc

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.06 11:48:00