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

Matplotlib多图合并:如何将多个多项式回归图整合为一张

合并多组多项式回归结果到单张图表

我写了一个多项式回归绘图函数,每次调用会生成独立图表。现在需要把多组子串对应的回归结果合并到同一张图里,流程要求如下:

  • 基于子串创建新DataFrame
  • 从DataFrame中提取x、y值
  • 划分训练测试集,筛选最优多项式阶数(最高10阶)并计算RMSE
  • 用最优阶数生成拟合函数
  • 按类别配色绘制原始数据点和拟合曲线
  • 标记拟合曲线在x=3、4、5处的交点

原调用方式会生成三张独立图:

ont = graph('ONT')
cmb = graph('CANHOU')
bcmfa = graph('BCMFA')

下面是修改后的代码,实现所有结果整合到单张图表:

import numpy as np
import matplotlib.pyplot as plt
from sklearn.linear_model import LinearRegression
from sklearn.preprocessing import PolynomialFeatures
from sklearn.metrics import mean_squared_error
from sklearn.model_selection import train_test_split

# 假设get_color_gradient是已实现的生成颜色渐变的函数
def get_color_gradient(start_color, end_color, num_points):
    start = np.array([int(start_color[i:i+2], 16) for i in (1,3,5)])
    end = np.array([int(end_color[i:i+2], 16) for i in (1,3,5)])
    gradient = np.linspace(start, end, num_points)
    return [f'#{int(r):02x}{int(g):02x}{int(b):02x}' for r,g,b in gradient]

def graph(number, ax):
    # 创建对应子串的DataFrame
    df = mat[mat['Security'].str.contains(number)]
    df = df.reset_index(drop=True).sort_values('Years')
    
    x = df['Years']
    y = df['Rate']
    
    # 筛选最优多项式阶数
    rmses = []
    degrees = range(1, 11)
    min_rmse, min_deg = 1e10, 0
    x_train, x_test, y_train, y_test = train_test_split(x, y, test_size=0.2)
    x_train = x_train.values.reshape(-1, 1)
    x_test = x_test.values.reshape(-1, 1)
    
    for deg in degrees:
        poly_features = PolynomialFeatures(degree=deg, include_bias=False)
        x_poly_train = poly_features.fit_transform(x_train)
        
        poly_reg = LinearRegression()
        poly_reg.fit(x_poly_train, y_train)
        
        # 用transform而非fit_transform避免测试集数据泄露
        x_poly_test = poly_features.transform(x_test)
        poly_predict = poly_reg.predict(x_poly_test)
        poly_rmse = np.sqrt(mean_squared_error(y_test, poly_predict))
        rmses.append(poly_rmse)
        
        if poly_rmse < min_rmse:
            min_rmse = poly_rmse
            min_deg = deg
            
    # 生成最优阶数的多项式拟合函数
    z = np.polyfit(x, y, min_deg)
    f = np.poly1d(z)
    
    # 根据传入的子串匹配对应的配色和名称
    gradient_blue = get_color_gradient('#003060', '#B9D9EB', len(x))
    gradient_red = get_color_gradient('#750000', '#F5D2D2', len(x))
    gradient_green = get_color_gradient('#1A4314', '#B2D2A4', len(x))
    gradient_purple = get_color_gradient('#4B0082', '#E6E6FA', len(x))
    
    if number == 'ONT':
        gradient = gradient_blue
        line_color = 'blue'
        name = 'Ontario'
    elif number == 'CANHOU':
        gradient = gradient_red
        line_color = 'red'
        name = 'CMB'
    elif number == 'BCMFA':
        gradient = gradient_green
        line_color = 'green'
        name = 'BCMFA'
    else:
        gradient = gradient_purple
        line_color = 'purple'
        name = 'Other'

    # 在传入的ax上绘制数据点和拟合曲线
    ax.scatter(x, y, color=gradient)
    ax.plot(x, f(x), label=f'{name} (阶数 {min_deg})', color=line_color)
    
    # 标记x=3、4、5处的交点并标注数值
    mat_pt = [3, 4, 5]
    mat_345 = np.interp(mat_pt, x, f(x))
    ax.scatter(mat_pt, mat_345, color=line_color, marker='*', s=200)
    for x_val, y_val in zip(mat_pt, mat_345):
        ax.annotate(f'{y_val:.2f}',
                    (x_val, y_val),
                    textcoords='offset points',
                    xytext=(15, -5),
                    ha='left',
                    color=line_color)

# 主程序:创建统一的图表轴,调用函数绘制所有数据
fig, ax = plt.subplots(figsize=(18, 8))

# 绘制三组数据
graph('ONT', ax)
graph('CANHOU', ax)
graph('BCMFA', ax)

# 添加通用的图表元素
ax.axvline(x=3, color='grey', linewidth=0.5)
ax.axvline(x=4, color='grey', linewidth=0.5)
ax.axvline(x=5, color='grey', linewidth=0.5)
ax.axhline(y=0, color='black', linewidth=0.5, linestyle='dashed')

ax.set_xlabel('Years', size=15)
ax.set_ylabel('Rate', size=15)
ax.set_title(f'Rate: {today_str}', size=20)
ax.legend(loc='upper left')

plt.show()

关键修改说明

  • 函数新增ax参数,指定统一的绘图坐标轴,避免每次创建独立图表
  • 修正原代码中bondtype未定义的问题,直接根据传入的number参数匹配类别信息
  • 将通用图表元素(垂直线、水平线、标题、坐标轴标签)移到函数外部统一设置,避免重复绘制
  • 调整测试集特征转换方式,使用transform而非fit_transform,防止数据泄露

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.29 16:07:08