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

如何基于返回单图的函数创建1×3布局的并排子图?

解决方法:让绘图函数支持指定子图轴

问题出在你的plot_learning_curve函数每次调用都会创建新的独立图形(通过plt.figure()),这和你预先创建的1×3子图网格完全无关——你赋值给axes[0][0]的操作根本无法把新生成的图绑定到已有的子图上,所以才会出现空网格+单独绘图的情况。

我们只需要给函数加一个可选的ax参数,让它可以指定在某个子图轴上绘图,同时保留原函数的独立绘图能力,不影响模块复用。

第一步:修改plot_learning_curve函数

def plot_learning_curve(estimator, X, y, ylim=None, cv=None, n_jobs=-1, 
                        train_sizes=np.linspace(.1, 1.0, 5), ax=None):
    """Generate a simple plot of the test and training learning curve"""
    # 如果没传入ax,就用当前轴(或创建新图形),保持原功能
    if ax is None:
        ax = plt.gca()
    
    ax.set_title(str(estimator).split('(')[0]+ " learning curves")
    if ylim is not None:
        ax.set_ylim(*ylim)
    ax.set_xlabel("Training examples")
    ax.set_ylabel("Score")
    
    train_sizes, train_scores, test_scores = learning_curve(
        estimator, X, y, cv=cv, n_jobs=n_jobs, train_sizes=train_sizes)
    
    train_scores_mean = np.mean(train_scores, axis=1)
    train_scores_std = np.std(train_scores, axis=1)
    test_scores_mean = np.mean(test_scores, axis=1)
    test_scores_std = np.std(test_scores, axis=1)
    
    ax.grid()
    ax.fill_between(train_sizes, train_scores_mean - train_scores_std,
                    train_scores_mean + train_scores_std, alpha=0.1, color="r")
    ax.fill_between(train_sizes, test_scores_mean - test_scores_std,
                    test_scores_mean + test_scores_std, alpha=0.1, color="g")
    ax.plot(train_sizes, train_scores_mean, 'o-', color="r", label="Training score")
    ax.plot(train_sizes, test_scores_mean, 'o-', color="g", label="Cross-validation score")
    ax.legend(loc="best")
    return ax

第二步:调用函数时传入子图轴

现在你可以把预先创建的子图轴传递给函数,让它在指定位置绘图:

fig, axes = plt.subplots(nrows=1, ncols=3, sharex="all", figsize=(15,5), squeeze=False)

# 给每个子图传入对应的ax参数
plot_learning_curve(tuned_clfs_vert_title2[0][0][1], Xs_train1, Y_train1, cv=skfold, ax=axes[0][0])
plot_learning_curve(tuned_clfs_vert_title2[0][1][1], Xs_train1, Y_train1, cv=skfold, ax=axes[0][1])
plot_learning_curve(tuned_clfs_vert_title2[0][2][1], Xs_train1, Y_train1, cv=skfold, ax=axes[0][2])

# 调整布局避免元素重叠
plt.tight_layout()
plt.show()

为什么这样可行?

  • 修改后的函数保持了向后兼容性:如果不传入ax参数,它的行为和原来完全一样,依然可以生成独立的单张图,不影响你把它作为模块单独使用。
  • 当传入ax时,函数会直接在你指定的子图轴上绘制内容,和预先创建的1×3网格完全绑定,不会生成额外的独立图形。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.29 07:05:53