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

