先说一个我经常遇到的场景:模型训练完了,SHAP分析也跑出来了,想用瀑布图解释单个样本的特征贡献,放进周报或者论文里。结果 shap.plots.waterfall 一执行,图是出来了,但四周多了个乌漆嘛黑的框架,坐标轴的刻度也乱乱地戳在外面。这时候很多人会顺手写一句 plt.box(False),然后眼睁睁看着那条边框纹丝不动。原因是,SHAP瀑布图的“边框”根本不是传统意义上的图框,而是 matplotlib 里的 axes spines、刻度线和文字占位互相叠加出的视觉污染,处理逻辑完全不是一行 .box(False) 能解决的。
这篇文章我打算把这件“小事”彻底讲透。内容围绕 SHAP 值瀑布图的定制,重点放在去除边框,但不止于去除边框。我会解释 SHAP 默认绘图函数为什么要用“隐藏样式”的方式帮你画图,也会给出三套从浅到深的解决方案:从直接擦除默认图边框,到从 Explanation 对象出发重绘一张属于自己的无水印瀑布图,再到把定制逻辑封装成可复用函数。无论你是刚入门的数据分析师,还是天天调模型的老手,应该都能找到可以抄走的那段代码。
1. 瀑布图默认样式的三个“隐形负担”
1.1 你以为的边框,其实是坐标轴脊线
先用一张图软件的角度来拆解。用户在 Jupyter 里看到 SHAP 瀑布图“四四方方的框”,通常由三种元素组成:
- 坐标轴的四个
spines,也就是 matplotlib 里上下左右四条脊线。默认情况下top和right可能没有,但如果样式模板给的是带框风格,四条线就全在; - 坐标轴上的刻度线
ticks,尤其是 Y 轴名字列表左右那些短横线,会让视觉上觉得“有一条没头没尾的线”; - 整个
Figure画布的背景色与坐标轴区域的背景色不一致,有些博客里还会叠加一层浅灰网格,最后导出 PNG 时看起来就像多个若隐若现的边框。
SHAP 瀑布图是从 matplotlib 的 Axes 画出来的,所以你可以用常规方式去清理。只不过它内部有很多自己的样式逻辑,导致很多人按标准姿势处理时发现无效,或者被它后续的绘制操作覆盖。
1.2 shap内部样式掩盖了哪些可控参数
你会发现,SHAP 大多数绘图函数在绘制过程中会调用自己的绘图上下文,比如:
python复制import shap
import matplotlib.pyplot as plt
shap.plots.waterfall(explanation_single, max_display=8)
一旦执行完这一句,当前 pyplot 里的一些全局配置可能已经变了。这意味着如果你在调用 shap.plots.waterfall 之前设置了自己的字体、坐标轴颜色或者 rcParams,有概率会被 SHAP 内部重置掉,至少是“部分覆盖”。这是很多同学折腾半天去不掉边框的最大原因:调整顺序错了。
我在实际项目里的习惯是,先把默认图画出来,再获取当前 Axes 去修改,而不是在画之前用 plt.style 或者 plt.rcParams 去和它较劲。SHAP 画完之后,plt.gca() 可以拿回当前活跃坐标轴,然后对 spines 动手。
1.3 为什么要关心这个:“能跑就行”和“能发表”的距离
有人会说,图能看就行,框不框的有什么关系。但如果你的输出要用于论文、PPT、报告,甚至是客户交付物,图是否干净直接影响专业观感。更关键的是,你后续可能要给瀑布图加标题、调整正负样本顺序、把多张图拼进一张画布里,这些操作如果建立在默认图的框线设置上,效果会非常不可控。
我见过不少次这种尴尬场面:分析师把两张瀑布图保存下来,用 PPT 的“删除背景”功能抠掉白底,结果因为图边上有半条脊线和刻度数字,怎么抠都有毛边。与其在展示端补救,不如在绘图端就把边框问题根治。
需要模型API调用? 免费领10W Token,多模型网关一键接入 Claude、DeepSeek 等主流模型。
2. 对默认瀑布图做“去框”手术
2.1 先拿到当前Axes再动手
最简单的方案,是允许 SHAP 先完成它的绘制,然后我们拿到它的 Axes 来清理。这里有个关键点,SHAP 的 shap.plots.waterfall 不一定直接返回一个 Axes,所以靠 plt.gca() 拿当前活跃轴是最稳妥的。
python复制import shap
import matplotlib.pyplot as plt
from sklearn.ensemble import RandomForestRegressor
from sklearn.datasets import load_diabetes
X, y = load_diabetes(return_X_y=True)
model = RandomForestRegressor(n_estimators=100, random_state=42)
model.fit(X, y)
explainer = shap.TreeExplainer(model)
shap_values = explainer(X) # 得到一个 Explanation 对象
# 选取第一个样本,先画出默认瀑布图
fig = plt.figure()
shap.plots.waterfall(shap_values[0], max_display=8, show=False)
# 此时 AXES 已经存在,直接拿当前轴
ax = plt.gca()
# 移除四个方向的脊线
for edge in ["top", "bottom", "left", "right"]:
ax.spines[edge].set_visible(False)
# 移除刻度线,但保留刻度文字(后面你大概率还要给特征排列显示)
ax.tick_params(which="both", length=0)
plt.tight_layout()
plt.show()
这段代码运行后,四条脊线会全部消失。注意我没有直接调用 plt.box(False),因为 plt.box 控制的是当前坐标轴是否显示边框,但它对“已经由 SHAP 用内部方式创建出来并且被塞进样式上下文”的轴,表现不一定如你所愿。直接遍历 spines 永远是更底层、更可控的姿势。
2.2 spines全部关闭后的残留来源
关掉 spines 之后你再截图,可能还会看到一些细碎框架,这些残留通常来自两个地方。
第一是刻度,特别是 Y 轴左侧的刻度数字和刻度线。瀑布图的 Y 轴通常是一组特征名,特征是文字标签,去掉刻度线后整张图会轻盈很多。但如果你的特征名本身是长字符串,默认图可能会在 Y 轴左侧做两次换行,看起来就像多出一列灰色小字,这不是边框,是标签排版问题。
第二是图例或注释文本的描边。matplotlib 的文字对象有时候会带白色描边,当图里其他元素被清掉后,文字的描边就会形成一圈若隐若现的白边,在某些深色背景截图下反而变成“白色边框”。处理办法是把相关文本对象的 set_path_effects 去掉,或者你自己重绘文字而不是依赖 SHAP 默认注释。
对于一张已经由官方函数画出来的默认图,能做到“移除四条边 + 关掉刻度线 + 使用无背景导出”,满足 90% 的去边框需求没有问题。
2.3 处理导出版边距:bbox_inches的妙用
另一种被误认为“去不掉边框”的情况,是保存图片时四周一圈白边。你截图时会自然带上画布完整信息,但在代码中手动保存,默认参数通常会在图像外留白。很多教程让你把 bbox_inches="tight" 写进 savefig,这一招的确管用:
python复制fig.savefig(
"waterfall_clean.png",
dpi=300,
bbox_inches="tight",
pad_inches=0.02,
facecolor="white",
)
这里 facecolor="white" 很重要。默认画布背景可能是透明,也可能带浅色,如果你想要一个纯白干净、无框线的输出,必须显式指定。bbox_inches="tight" 会重新计算图内容的包围盒,把多余白边剔除。pad_inches=0.02 是预留一点点边缘,避免某些文本被裁切。
如果你讨厌的是保存后四边有“框线式”的白边,这招比去 spines 还直接。
3. 从Explanation出发重绘无边框瀑布图
3.1 提取base_values、values、feature_names
直接清理默认图虽然快,但也有局限。很多定制需求,比如“给负向贡献的文本加粗”“在每个特征条上标注原始特征值”“把数字标签放到条内部”,默认函数并不直接支持。这时候就得从 SHAP 的 Explanation 对象里把原料拉出来,自己画一张。
Explanation 对象的结构,拆开来看并不复杂:
base_values:基线预测值,也就是全部样本的平均预测值;values:每个特征的 SHAP 贡献值数组;data:当前样本对应特征的原始值,绘图时可以当作文本展示;feature_names:特征名列表。
代码上这么取:
python复制import numpy as np
import shap
single = shap_values[0] # Explanation 对象里的一行
base_value = float(single.base_values)
shap_vals = np.asarray(single.values).ravel()
feature_names = [str(n) for n in single.feature_names]
raw_values = np.asarray(single.data).ravel()
# 按贡献绝对值排序,方便取 Top N
order = np.argsort(np.abs(shap_vals))[::-1]
max_display = 8
order = order[:max_display]
feature_names = [feature_names[i] for i in order]
shap_vals = shap_vals[order]
raw_values = raw_values[order]
从 Explanation 对象着手的好处是,后续所有视觉元素都归你管,再也没有 SHAP 的隐形样式来打扰。这也是我能彻底解决“边框”问题的根本手段:既然轴是我自己创建的,我不画 spines,它就不存在。
3.2 水平瀑布图绘制逻辑和核心代码
瀑布图的含义是:从基线预测值出发,每个特征的贡献一层层累加,最终到达模型对当前样本的预测值。所以画出“每一段的起点和终点”是关键。
python复制import matplotlib.pyplot as plt
def plot_custom_waterfall(
base_value,
shap_vals,
feature_names,
raw_values,
max_display=8,
figsize=(10, 6),
title=None,
save_path=None,
):
order = np.argsort(np.abs(shap_vals))[::-1][:max_display]
names = [feature_names[i] if i < len(feature_names) else f"feature_{i}" for i in order]
vals = shap_vals[order]
# 累计位置
cumulative = base_value + np.r_[0, np.cumsum(vals)]
fig, ax = plt.subplots(figsize=figsize)
ax.set_facecolor("white")
ys = np.arange(len(vals))[::-1] # 最重要的特征放最上面
for i, y in enumerate(ys):
start = cumulative[i]
end = cumulative[i + 1]
color = "#d62728" if end - start >= 0 else "#1f77b4"
# 用粗横线表示一段贡献
ax.hlines(
y=y,
xmin=start,
xmax=end,
linewidth=10,
solid_capstyle="round",
color=color,
zorder=3,
)
# 在条中间标注贡献值
text_pos = start + (end - start) / 2
ax.text(
text_pos,
y + 0.18,
f"{end - start:+.3f}",
ha="center",
va="bottom",
fontsize=9,
color="#333333",
)
# 左边放当前样本的原始特征值
ax.text(
0.01,
y,
f"{names[i]} = {raw_values[i]:.3f}",
transform=ax.get_yaxis_transform(),
ha="left",
va="center",
fontsize=10,
color="#222222",
)
# 画一条基线竖直线
ax.axvline(base_value, color="#999999", linestyle="--", linewidth=1)
# 不画任何框线,这里从源头解决问题
for spine in ax.spines.values():
spine.set_visible(False)
ax.set_yticks([])
ax.tick_params(left=False)
ax.set_xlabel("prediction = %.3f" % cumulative[-1], fontsize=11)
if title:
ax.set_title(title, fontsize=13, loc="left")
if save_path:
fig.savefig(save_path, dpi=300, bbox_inches="tight", pad_inches=0.02, facecolor="white")
return fig, ax
这段代码里我已经彻底不画 Y 轴刻度,而是把特征名和原始值直接作为文本放进左侧。这样做与 SHAP 默认图不同,但信息量更大,因为 SHAP 官方瀑布图有时会把“特征名”和“原始值”拆到两个图层,导致很多人导出 SVG 再编辑时找不到文字从哪来。
3.3 样式细节:将特征值作为第二信息通道
SHAP 值瀑布图最怕一个误区:只展示 SHAP 贡献,却不展示该特征在当前样本里的实际取值。如果贡献值是负的,读者会想知道“到底这个特征是处于低值还是高值导致贡献为负”。所以你重绘时,最好在每条特征条形旁边带上原始值。
默认的 Shap 图虽然也显示特征值,但它用的是特征名 + 灰色的原始数值,排布比较紧密。自己做的时候,我会把特征值放在左侧文本里,用类似:
python复制age = 36
这样的一条文本直接说明当前样本该特征的真实取值。这种用法在做客户流失分析、信贷风控样本解释时特别有用。
3.4 文字描边、箭头线和颜色自定义
自己重绘时,有一个关于“边框感”的隐藏坑在文字层。matplotlib 的 text 默认没有描边,但当你给文字设置了 path_effects 这样的效果,比如:
python复制import matplotlib.patheffects as path_effects
text = ax.text(0, 0, "example", fontsize=12)
text.set_path_effects([
path_effects.withStroke(linewidth=3, foreground="white")
])
在视觉上每个字都会带上白色描边。如果你之后把图放到深色背景或者透明底上,这些白边会形成非常明显的“伪边框”。更隐蔽的是,有些 SHAP 版本内部给标题和刻度文字设置了类似的描边,清空 spines 后才发现文字旁边还是脏脏的。
所以我的建议是:如果追求零边框洁癖效果,干脆完全自定义所有文字,不要依赖 SHAP 内部的 text 对象。你控制每一个字体、字号和描边,才能保证导出的图是真正干净的。
4. 多组样本对比与版本差异坑
4.1 把两张瀑布图并排且不透出框线的布置
现实生活中只解释一个样本的情况并不多,通常分析师会希望把两个对比样本放在同一张图里,比如“流失用户”和“留存用户”的特征贡献并排比较。这种需求用 SHAP 默认瀑布图很难做,因为它一次只处理一条 Explanation,并且内部可能会重新调整画布比例。
我的做法是,使用自定义绘图函数,并封装一个支持传 ax 参数的内核。把刚才 plot_custom_waterfall 里的绘图逻辑抽出来,允许外部指定 ax:
python复制def draw_waterfall_on_ax(ax, single, max_display=8):
base_value = float(single.base_values)
values = np.asarray(single.values).ravel()
feature_names = [str(n) for n in single.feature_names]
order = np.argsort(np.abs(values))[::-1][:max_display]
vals = values[order]
names = [feature_names[i] for i in order]
cumulative = base_value + np.r_[0, np.cumsum(vals)]
ys = np.arange(len(vals))[::-1]
for i, y in enumerate(ys):
start = cumulative[i]
end = cumulative[i + 1]
color = "#d62728" if end - start >= 0 else "#1f77b4"
ax.hlines(y=y, xmin=start, xmax=end, linewidth=8,
solid_capstyle="round", color=color)
ax.text(start + (end - start) / 2, y + 0.2,
f"{end - start:+.3f}", ha="center", fontsize=8)
ax.axvline(base_value, color="#999999", linestyle="--", linewidth=1)
ax.set_yticks([])
for spine in ax.spines.values():
spine.set_visible(False)
# 创建多个子图
fig, axes = plt.subplots(1, 2, figsize=(15, 6))
for ax, single_instance in zip(axes, [shap_values[0], shap_values[1]]):
draw_waterfall_on_ax(ax, single_instance)
ax.set_title(f"sample with prediction = {single_instance.base_values + single_instance.values.sum():.3f}", loc="left")
axes[0].set_xlabel("churn = 0")
axes[1].set_xlabel("churn = 1")
plt.tight_layout()
plt.show()
子图模式下,每个 ax 都有自己的坐标轴背景和 spines。只要 draw_waterfall_on_ax 里统一关掉 spines,并按需隐藏刻度,多张图拼在一起时框线就不会互相干扰。
4.2 使用shap 0.45前后API变化
SHAP 的版本迭代比较快,API 变化在可视化函数上尤其明显。最常见的是旧版代码:
python复制shap.waterfall_plot(shap_values[0])
新版推荐的是:
python复制shap.plots.waterfall(shap_values[0])
新版更面向 Explanation 对象,返回的数据结构也更规范。如果你写代码时直接用官方示例,需要先确认自己装的是不是新版。版本差异还会影响 shap_values 到底是 numpy 数组还是 Explanation 对象。
遇到这种情况,我一般先跑:
python复制print(type(shap_values))
如果是 shap.Explanation,那么 shap_values[i] 的属性访问方式才安全。如果用旧索引方式取多维数组再接 base_values,很容易碰到 Tuple 索引越界。
4.3 长特征名和小图幅的隐性边框
多图对比时另一个隐蔽问题是:特征名太长,文字会溢出到子图边界之外,视觉上像多了一条横线。比如中文特征名或带单位的长字符串,如果不在画布内做换行或截断,就会越过坐标轴区域,直接叠在另一个子图上。
我通常会做三件事:
- 限制
feature_names的最大长度,超过 20 个字符就用省略号; - 用
textwrap对长特征名做换行,比如每 15 个字符断行; - 调高
figsize的高度,或减少max_display,保证横向空间充足。
很多“去除边框”的抱怨,其实根因是标签溢出后形成的误视觉。文字一旦出了绘制区域,和边框线叠在一起,你怎么擦 spines 都没用。
5. 沉淀一个可直接调用的定制函数
5.1 完整实现
把以上经验沉淀下来,我平时会维护一个 plot_waterfall_clean 函数,功能上覆盖“提取信息、去框线、标原始特征值、可保存、可拼图”。核心逻辑如下:
python复制import textwrap
def plot_waterfall_clean(
shap_exp,
max_display=10,
figsize=(11, 6),
save_path=None,
title=None,
ax=None,
):
"""把 SHAP Explanation 画成无边框定制瀑布图。"""
single = shap_exp
base_value = float(single.base_values)
values = np.asarray(single.values).ravel()
feature_names = [str(n) for n in single.feature_names]
raw_data = np.asarray(single.data).ravel() if single.data is not None else None
order = np.argsort(np.abs(values))[::-1][:max_display]
vals = values[order]
names = [feature_names[i] for i in order]
raw_vals = [raw_data[i] if raw_data is not None else None for i in order]
if ax is None:
fig, ax = plt.subplots(figsize=figsize)
else:
fig = ax.figure
cumulative = base_value + np.r_[0, np.cumsum(vals)]
ys = np.arange(len(vals))[::-1]
# 绘制非显著特征的透明指示条
hidden_contribution = values.sum() - vals.sum()
if len(vals) < len(values):
ax.hlines(
y=-1,
xmin=cumulative[0],
xmax=cumulative[-1],
linewidth=3,
color="#cccccc",
linestyle=":",
)
ax.text(
cumulative[0],
-0.8,
f"... {len(values) - len(vals)} other features sum = {hidden_contribution:+.3f}",
fontsize=9,
color="#777777",
)
for i, y in enumerate(ys):
start = cumulative[i]
end = cumulative[i + 1]
color = "#d62728" if end - start >= 0 else "#1f77b4"
ax.hlines(y=y, xmin=start, xmax=end, linewidth=10,
solid_capstyle="round", color=color)
# 贡献值标注放段内偏右
pos = end if end > start else start
ax.text(pos, y + 0.15, f"{end - start:+.3f}",
fontsize=9, va="bottom", color="#444444")
# 左侧显示特征名和原始值
if raw_vals[i] is not None:
label = f"{names[i]} = {raw_vals[i]:.3f}"
else:
label = f"{names[i]}"
wrapped = textwrap.fill(label, width=30)
ax.text(-0.02, y, wrapped, transform=ax.get_yaxis_transform(),
ha="right", va="center", fontsize=10, color="#111111")
ax.axvline(base_value, color="#888888", linestyle="--", linewidth=1)
ax.set_xlabel(f"prediction = {cumulative[-1]:.3f}", fontsize=10)
ax.set_yticks([])
ax.tick_params(length=0)
for spine in ax.spines.values():
spine.set_visible(False)
if title:
ax.set_title(title, loc="left", fontsize=13)
ax.set_facecolor("white")
if save_path:
fig.savefig(save_path, dpi=300, bbox_inches="tight", pad_inches=0.02, facecolor="white")
return ax
这段封装把几个容易忽视的边界情况都处理了:
raw_data可能为空,原始值展示要兜底;- 显示 Top 10 时,累计路径达不到最终预测值,需要标注还有多少个特征被隐藏;
- 长文本在左轴内侧展示,不额外占用右边空间。
5.2 调用示例和进一步扩展
调用方式非常直接:
python复制shap_exp = shap_values[0]
plot_waterfall_clean(
shap_exp,
max_display=8,
save_path="waterfall_sample.png",
title="Customer #1001 - feature attribution",
)
如果想把这张图输出到论文排版环境,我建议导出 SVG 格式而不是透明 PNG。SVG 是矢量格式,任何文本处理软件都能继续编辑,不需要担心“框线”或者像素问题。
python复制plot_waterfall_clean(
shap_values[1],
max_display=8,
save_path="waterfall_sample.svg",
)
SVG 文件里,如果某个文本还需要在 Illustrator 里手工修改,你可以再设置字体为通用字体,不要用 Jupyter 默认字体,否则换电脑打开时可能出现替代字体错位。
最后再分享一个我的真实体会:很多定制问题,最初看起来是“边框线去除不了”这种视觉细节,但往深了挖,往往会发现是对绘图对象的不了解。无论是 SHAP 还是其他可视化库,默认画图方便,可是当你要进入定制阶段,最好的办法永远是放弃“擦除默认图上的东西”,而是“只新建一张由你完全控制的图”。前者是不断打补丁,后者是从根上解决问题。
我用这个思路处理 SHAP 瀑布图之后,不仅是边框,包括颜色、字号、标签位置、导出格式,都变得非常顺手。希望这套方法也能帮你把图做得干净、专业。
