08_polynomialFitting.py - 多项式拟合与插值完全指南

学习路径第 8 步 (共 10 步) | 难度:中高级

概述

polyfit / poly1d 基础操作出发,深入理解过拟合与欠拟合,学习加权拟合、多项式微积分运算,以及插值 vs 拟合的本质区别

学习目标

  • 掌握 np.polyfit()np.poly1d() 进行多项式拟合
  • 理解多项式的加减乘除、求导、积分运算
  • 通过可视化理解过拟合欠拟合
  • 学会使用加权拟合处理不等精度数据
  • 了解拉格朗日插值 / 样条插值的适用场景

核心内容 (7 个模块)

模块 核心知识点
1. polyfit & poly1d 基础拟合、多项式求值
2. 多项式运算 加减乘除、求导(deriv)、积分(integ)
3. 曲线回归实战 非线性关系建模、真实数据拟合
4. 过拟合与欠拟合 不同阶数模型的可视化对比
5. 加权拟合 处理异方差数据的权重分配策略
6. 插值方法 Lagrange 插值 / 样条插值(spline)
7. 根与极值点 多项式求根、求极值

code

#!/usr/bin/env python
# -*- coding: utf-8 -*-

"""
=====================================
NumPy 多项式拟合与插值完全指南 (Polynomial Fitting)
=====================================

本案例系统介绍 NumPy 的多项式处理能力:

1. polyfit / poly1d — 多项式拟合
2. 多项式的运算 (加减乘除、微积分)
3. 曲线回归实战 (非线性关系建模)
4. 多项式插值 vs 拟合的区别
5. 过拟合与欠拟合的可视化理解
6. 实战案例: 温度数据建模 / 趋势预测

【核心价值】
  多项式拟合是最基础的回归分析方法之一,
  是理解更复杂模型(如神经网络、核方法)的重要起点。

作者:bloxed
"""

import numpy as np


def separator(title):
    print(f"\n{'='*60}")
    print(f"  {title}")
    print('='*60)


# ============================================================
# 第一部分:多项式拟合基础
# ============================================================
separator("一、多项式拟合基础 (polyfit + polyval)")

print("""
【任务】给定一组数据点 (xi, yi),找到一个多项式
       p(x) = c_n·x^n + ... + c_1·x + c_0
       使其最好地拟合这些点。

【核心函数】
  np.polyfit(x, y, deg)  → 返回系数 [c_n, ..., c_1, c_0] (降幂排列)
  np.polyval(p, x)       → 计算多项式在 x 处的值
  np.poly1d(coefficients) → 创建多项式对象 (方便调用和显示)
""")

# 生成带噪声的非线性数据
rng = np.random.default_rng(42)
n = 40
x_data = np.linspace(-3, 3, n)
# 真实函数: y = 2x^2 - x + 1 + 噪声
y_true = 2 * x_data**2 - x_data + 1
y_data = y_true + rng.normal(0, 1.5, size=n)

print(f"数据点数: {n}")
print(f"x范围: [{x_data.min()}, {x_data.max()}]")
print(f"真实关系: y = 2x^2 - x + 1 + noise")

# 尝试不同阶数的拟合
degrees = [1, 2, 3, 8]
print(f"\n{'阶数':>4} {'系数 (降幂)':<35} {'MSE':>10}")
print("-" * 55)

results = {}
for deg in degrees:
    coeffs = np.polyfit(x_data, y_data, deg)
    y_pred = np.polyval(coeffs, x_data)
    mse = np.mean((y_data - y_pred)**2)
    
    coeff_str = ", ".join([f"{c:.3f}" for c in coeffs])
    results[deg] = (coeffs, y_pred, mse)
    print(f"{deg:>4}  [{coeff_str:<33}] {mse:>10.3f}")

# 展示 poly1d 对象
p2 = np.poly1d(results[2][0])  # 2阶多项式对象
print(f"\n2阶拟合的多项式对象:")
print(f"  p(x) = \n{p2}")
print(f"  p(0) = {p2(0):.3f}")
print(f"  p(1) = {p2(1):.3f}")


# ============================================================
# 第二部分:多项式对象的完整操作
# ============================================================
separator("二、多项式对象 (poly1d) 完整操作")

print("""
np.poly1d 不仅存储系数, 还提供丰富的多项式操作:
""")
# 定义两个多项式
p_a = np.poly1d([1, -3, 2])    # x^2 - 3x + 2 = (x-1)(x-2)
p_b = np.poly1d([1, 1])         # x + 1

print(f"p_a(x) = {p_a}")       # 显示多项式
print(f"p_b(x) = {p_b}")

print(f"\n--- 基本运算 ---")
print(f"p_a + p_b  = {p_a + p_b}")
print(f"p_a - p_b  = {p_a - p_b}")
print(f"p_a * p_b  = {p_a * p_b}")
print(f"p_a / p_b  = 除法得到 (商, 余数): {np.polydiv(p_a.coeffs, p_b.coeffs)}")

print(f"\n--- 求值 ---")
print(f"p_a(0) = {p_a(0)}")       # 单点求值
print(f"p_a([0,1,2]) = {p_a([0,1,2])}")  # 多点批量求值

print(f"\n--- 微积分 ---")
print(f"p_a 的导数: {np.polyder(p_a)}")
print(f"p_a 的不定积分: {np.polyint(p_a)}")

# 定积分
integral_val = np.polyint(p_a)
area = integral_val(1) - integral_val(0)
print(f"∫[0,1] p_a(x)dx = {area:.3f}")

print(f"\n--- 根 ---")
roots = np.roots(p_a.coeffs)
print(f"p_a(x)=0 的根: x = {roots}")
print(f"验证: p_a({roots[0]:.4f}) = {p_a(roots[0]):.6f}")


# ============================================================
# 第三部分:过拟合 vs 欠拟合
# ============================================================
separator("三、过拟合 vs 欠拟合 (模型复杂度的艺术)")

print("""
【关键概念】

  欠拟合 (Underfitting): 模型太简单,无法捕捉数据规律
    → 高偏差 (high bias)
    
  适当拟合 (Good fit): 模型复杂度适中
    
  过拟合 (Overfitting): 模型太复杂,记住了噪声
    → 高方差 (high variance), 泛化能力差
""")

# 重新生成数据
np_rng = np.random.default_rng(77)
n_train = 25
x_train = np.linspace(0, 2*np.pi, n_train)
# 真实函数: y = sin(x) + 小噪声
y_train = np.sin(x_train) + np_rng.normal(0, 0.2, n_train)

# 测试集
x_test = np.linspace(0, 2*np.pi, 100)
y_test_true = np.sin(x_test)

print(f"训练数据: {n_train} 个点 (sin曲线 + 噪声)")
print(f"测试数据: 100 个点 (纯净 sin 曲线)\n")

test_mses = {}
train_mses = {}

for degree in [1, 2, 3, 5, 10, 15]:
    coeffs = np.polyfit(x_train, y_train, degree)
    p = np.poly1d(coeffs)
    
    y_train_pred = p(x_train)
    y_test_pred = p(x_test)
    
    train_mse = np.mean((y_train - y_train_pred)**2)
    test_mse = np.mean((y_test_true - y_test_pred)**2)
    
    train_mses[degree] = train_mse
    test_mses[degree] = test_mse

print(f"{'阶数':>4} {'训练MSE':>12} {'测试MSE':>12} {'状态':<16}")
print("-" * 48)
for d in train_mses.keys():
    train_err = train_mses[d]
    test_err = test_mses[d]
    if d <= 2:
        status = "欠拟合"
    elif d >= 8:
        status = "过拟合!"
    else:
        status = "良好"
    print(f"{d:>4} {train_err:>12.4f} {test_err:>12.4f} {status}")

print(f"""
[观察]
  低阶 (deg=1,2):  训练和测试都差 → 欠拟合
  中等阶 (deg=3-5): 测试误差最低 → 最佳泛化
  高阶 (deg=10,15): 训练误差≈0 但测试误差↑ → 过拟合!

[原则] Occam's Razor (奥卡姆剃刀):
  选择能合理解释数据的 **最简单** 模型
""")


# ============================================================
# 第四部分:加权多项式拟合
# ============================================================
separator("四、加权多项式拟合 (Weighted Polyfit)")

print("""
【何时需要权重?】

  不是所有数据点的可信度相同:
  - 近期数据比历史数据更重要
  - 精密仪器的测量比粗略测量更可靠
  - 中间区域的测量比边界区域更准确

  np.polyfit 支持 w 参数指定每个点的权重!
""")

# 模拟: 边界处测量误差更大
x_w = np.linspace(0, 10, 30)
y_w = 0.5 * x_w**2 + 2 * x_w + 3 + np.random.RandomState(45).normal(0, 3, 30)
# 边界处的测量噪声更大 (中间高置信度)
weights = np.exp(-((x_w - x_w.mean())**2) / (2 * (x_w.std())**2))

# 不加权和加权的对比
coeffs_unw = np.polyfit(x_w, y_w, deg=2)
coeffs_wtd = np.polyfit(x_w, y_w, deg=2, w=weights)

print(f"不加权拟合: y = {coeffs_unw[0]:.4f}x^2 + {coeffs_unw[1]:.4f}x + {coeffs_unw[2]:.4f}")
print(f"加权拟合:   y = {coeffs_wtd[0]:.4f}x^2 + {coeffs_wtd[1]:.4f}x + {coeffs_wtd[2]:.4f}")
print(f"真实关系:   y = 0.5x^2 + 2x + 3")
print(f"\n[!] 加权后中间高置信度区域的拟合更接近真实参数!")


# ============================================================
# 第五部分:多项式插值
# ============================================================
separator("五、多项式拟合 vs 插值 (关键区别)")

print("""
【拟合 (Fitting)】
  找一条曲线 **逼近** 所有点
  点数 >> 参数数量
  允许一定误差 (最小化残差平方和)
  应用: 回归分析、趋势预测

【插值 (Interpolation)】
  找一条曲线 **穿过** 所有点
  点数 == 参数数量 (或使用分段低阶插值)
  严格经过每个数据点
  应用: 数据填补、平滑、重采样
""")

# 插值演示
x_interp_nodes = np.array([0, 1, 2, 3, 4])
y_interp_nodes = np.array([0, 1, 0.5, 2, 1])

# 拉格朗日插值 (n个点 → n-1阶多项式, 严格通过所有点)
interp_coeffs = np.polyfit(x_interp_nodes, y_interp_nodes, deg=len(x_interp_nodes)-1)
p_interp = np.poly1d(interp_coeffs)

# 在更多点上评估
x_fine = np.linspace(0, 4, 50)
y_interp_fine = p_interp(x_fine)

print(f"插值节点: {list(zip(x_interp_nodes, y_interp_nodes))}")
print(f"插值多项式阶数: {len(x_interp_nodes)-1}")
print(f"\n验证插值点是否严格通过:")
for xv, yv in zip(x_interp_nodes, y_interp_nodes):
    calc_y = p_interp(xv)
    match = "OK" if abs(calc_y - yv) < 1e-10 else "FAIL!"
    print(f"  x={xv}, 真实y={yv}, 插值y={calc_y:.6f}{match}")

print(f"""
[!] 注意: 高阶插值可能产生龙格现象 (Runge's Phenomenon)
  在区间两端剧烈振荡! 这就是为什么实际中常使用:
  - 分段低阶插值 (如样条 spline)
  - 或拟合而非插值
""")


# ============================================================
# 第六部分:实战 —— 温度数据分析与趋势预测
# ============================================================
separator("六、实战: 月度温度数据建模与预测")

print("""
【场景】某城市过去12个月的月均温度数据
  任务:
  1. 用多项式拟合年度温度周期
  2. 评估拟合质量
  3. 预测未来几个月的温度趋势
""")

months = np.arange(1, 13)
# 模拟北半球城市: 冬冷夏热 (余弦波形 + 趋势)
base_temp = 15  # 年均温
amplitude = 12  # 季节波动幅度
phase_shift = np.pi / 6  # 7月最热
trend = 0.02  # 微弱升温趋势

temperature = base_temp - amplitude * np.cos(2*np.pi*(months - 7)/12) + trend * months
temperature += np.random.RandomState(88).normal(0, 1.5, 12)
temperature = np.round(temperature, 1)

print("月度平均温度 (°C):")
print(f"  {'月份':>4}", end="")
for m in months:
    print(f"{m:>6}", end="")
print()
print(f"  {'温度':>4}", end="")
for t in temperature:
    print(f"{t:>6.1f}", end="")
print()

# 用傅里叶风格的拟合 (cos + sin 组合, 等效于相位调整)
# 这里用普通多项式拟合作为示例
for deg in [2, 4, 6]:
    coeffs = np.polyfit(months, temperature, deg)
    p = np.poly1d(coeffs)
    fitted = p(months)
    rmse = np.sqrt(np.mean((temperature - fitted)**2))
    print(f"\n  {deg}阶拟合 RMSE: {rmse:.2f}°C")

# 选择4阶进行预测
best_deg = 4
best_p = np.poly1d(np.polyfit(months, temperature, best_deg))

future_months = np.arange(13, 19)  # 未来6个月
predicted_temp = best_p(future_months)

print(f"\n--- 使用 {best_deg} 阶多项式预测 ---")
print(f"  {'月份':>6} {'预测温度':>10}")
for m, t in zip(future_months, predicted_temp):
    label = f"{m}月(明年)" if m <= 12 else f"{m-12}月(后年)"
    print(f"  {label:>6} {t:>9.1f}°C")

print(f"""
[注意] 多项式外推的风险:
  多项式在训练范围之外的行为往往不稳定!
  对于周期性数据, 使用傅里叶级数或季节性ARIMA更合适。
  此处仅为演示多项式拟合流程。
""")


# ============================================================
# 第七部分:多元多项式回归 (简述)
# ============================================================
separator("七、扩展: 多元多项式回归 (特征工程)")

print("""
【概念】当因变量受多个自变量影响时:

  一元: y = f(x)          → np.polyfit(x, y, deg)
  多元: y = f(x1, x2, ...) → 需要手动构造设计矩阵

【示例】二元二次: y = c0 + c1·x1 + c2·x2 + c3·x1² + c4·x2² + c5·x1·x2
  设计矩阵 X = [1, x1, x2, x1², x2², x1·x2]
""")

# 二元数据
n_obs = 50
rng_multi = np.random.default_rng(55)
x1_data = rng_multi.uniform(-1, 1, n_obs)
x2_data = rng_multi.uniform(-1, 1, n_obs)
# 真实关系: y = 1 + 2x1 + 3x2 + x1² + 0.5x2² + x1*x2 + noise
y_multi = (1 + 2*x1_data + 3*x2_data + x1_data**2 
           + 0.5*x2_data**2 + x1_data*x2_data 
           + rng_multi.normal(0, 0.3, n_obs))

# 构造设计矩阵 (特征工程)
X_design = np.column_stack([
    np.ones(n_obs),        # 常数项
    x1_data,               # x1
    x2_data,               # x2
    x1_data**2,            # x1²
    x2_data**2,            # x2²
    x1_data * x2_data,     # x1·x2 交叉项
])

# 最小二乘求解
coeffs_multi, residual, _, _ = np.linalg.lstsq(X_design, y_multi, rcond=None)

terms = ['常数项', 'x1', 'x2', 'x1²', 'x2²', 'x1·x2']
true_coeffs = [1, 2, 3, 1, 0.5, 1]

print(f"  {'项':<8} {'拟合值':>8} {'真实值':>8} {'误差':>10}")
print("-" * 36)
for term, fit, true_val in zip(terms, coeffs_multi, true_coeffs):
    err = fit - true_val
    print(f"  {term:<8} {fit:>8.3f} {true_val:>8.1f} {err:>+10.3f}")

rmse = np.sqrt(residual / n_obs) if len(residual) > 0 else 0
print(f"\n  RMSE: {rmse:.3f}")
print(f"  R² = {1 - residual / np.sum((y_multi - y_multi.mean())**2):.4f}")


# ============================================================
# 总结
# ============================================================
separator("总结: 多项式拟合速查")

summary = """
+------------------------------------------------------------+
|  NumPy Polynomial Functions                                |
+------------------------------------------------------------+
|                                                            |
|  [拟合与求值]                                               |
|  np.polyfit(x, y, deg)   → 多项式系数 (降幂)              |
|  np.polyval(p, x)        → 计算多项式值                    |
|  np.poly1d(c)             → 创建多项式对象                  |
|                                                            |
|  [多项式运算]                                               |
|  np.polyadd(p1, p2)      → 多项式相加                     |
|  np.polysub(p1, p2)      → 多项式相减                     |
|  np.polymul(p1, p2)      → 多项式相乘                     |
|  np.polydiv(p1, p2)      → 多项式除法 (商,余数)           |
|  np.polyder(p)           → 求导                           |
|  np.polyint(p)           → 不定积分                       |
|                                                            |
|  [根与工具]                                                |
|  np.roots(c)              → 求多项式的根                   |
|  np.polyfromroots(r)     → 由根构建多项式                  |
|  np.polysub(x, y)        → 减法                           |
|                                                            |
|  [拟合参数]                                                 |
|  deg        → 多项式阶数                                   |
|  w          → 各点权重                                     |
|  cov=True   → 返回协方差矩阵 (估计参数不确定性)            |
|  full=True  → 返回额外诊断信息                             |
|                                                            |
|  [模型选择原则]                                             |
|  • 从简单模型开始, 逐步增加复杂度                          |
|  • 用验证集/交叉验证评估泛化性能                           |
|  •警惕过拟合: 训练好但测试差                               |
|  • 周期数据考虑傅里叶而非纯多项式                          |
+------------------------------------------------------------+
"""
print(summary)

print("\n运行完毕!")
Logo

汇聚全球AI编程工具,助力开发者即刻编程。

更多推荐