如何在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
相关产品推荐
相关产品推荐

