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
相关产品推荐
相关产品推荐

