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

如何在自定义Matplotlib函数中实现多图并排展示?

解决偏依赖图多图并排展示的问题

错误原因

当用plt.subplots(1, n)创建1行多列布局时,返回的axes是numpy数组,而非单个轴对象。如果你的自定义plot_pdp函数直接对这个数组调用grid()、plot()等轴方法,就会触发AttributeError——数组本身没有这些属性,只有数组里的单个轴对象才具备。

解决方案

核心是让自定义函数支持接收指定的子图轴对象,循环时为每个特征列分配对应的子图轴。

步骤1:修改自定义plot_pdp函数

新增ax参数,默认值设为None,未传入时自动创建新轴;传入则使用指定轴绘图:

from sklearn.inspection import partial_dependence
import matplotlib.pyplot as plt

def plot_pdp(model, X, feature_name, ax=None):
    # 未传入轴则创建新轴
    if ax is None:
        ax = plt.gca()
    
    # 计算偏依赖结果
    pdp_results = partial_dependence(model, X=X, features=[feature_name])
    feature_values = pdp_results['values'][0]
    pdp_values = pdp_results['average'][0]
    
    # 在指定轴上绘制偏依赖图
    ax.plot(feature_values, pdp_values, linewidth=2)
    ax.set_title(f"Partial Dependence: {feature_name}")
    ax.set_xlabel(feature_name)
    ax.set_ylabel("Partial Dependence")
    ax.grid(True)  # 操作单个轴对象,不会触发数组属性错误

步骤2:创建多列布局并循环绘图

根据特征数量创建1行多列的子图,遍历特征列和对应轴对象调用修改后的函数:

import numpy as np

# 获取数据集所有特征列名
feature_names = X.columns
n_features = len(feature_names)

# 创建1行n_features列的子图,调整尺寸避免拥挤
fig, axes = plt.subplots(1, n_features, figsize=(5*n_features, 4))

# 统一处理单/多特征场景:将axes转为一维数组
axes = np.ravel(axes)

# 遍历特征并绘制到对应子图
for idx, feature in enumerate(feature_names):
    plot_pdp(model, X, feature, ax=axes[idx])

# 调整子图间距,避免标签、标题重叠
plt.tight_layout()
plt.show()

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.20 18:18:50