Skip to content

回归算法概览

回归分析(Regression analysis)是一种基础的统计和机器学习技术,用于建模因变量(目标或结果)与一个或多个自变量(特征或预测变量)之间的关系。在监督学习(Supervised learning)中,回归的主要目标是预测一个连续数值。

与分类(Classification,预测离散类别)不同,回归预测的是诸如价格、温度、身高、销售额或传感器读数等连续量。模型从训练数据中学习输入特征(Features)与相应的连续输出值之间的关联。

令 X 表示输入特征(自变量),y 表示连续目标变量(因变量)。回归模型旨在学习一个函数 f,使得 y ≈ f(X)。

根据自变量的数量,回归模型可大致分为:

  • 简单线性回归(Simple Linear Regression): 使用一条直线(y = mx + c)建模单个自变量(特征)与连续因变量(目标变量)之间的关系。
  • 多元线性回归(Multiple Linear Regression): 使用一个线性方程(y = b₀ + b₁x₁ + b₂x₂ + ... + bₚxₚ)建模多个自变量(特征)与连续因变量之间的关系。

除了线性模型,还存在各种其他回归算法,以捕捉更复杂的非线性关系。

使用 Scikit-learn 在 Python 中构建回归器

Section titled “使用 Scikit-learn 在 Python 中构建回归器”

与分类器类似,Scikit-learn 为构建回归模型提供了一致的接口。让我们构建一个基本的简单线性回归器。

典型步骤包括:

  1. 导入必要的库。
  2. 加载或生成数据集。
  3. 将数据拆分为训练集和测试集。
  4. 选择并实例化回归模型。
  5. 在训练数据上训练(拟合)模型。
  6. 在测试数据上进行预测。
  7. 评估模型性能。
# Step 1: Import necessary libraries
import numpy as np
import matplotlib.pyplot as plt
import seaborn as sns
from sklearn.model_selection import train_test_split
from sklearn.linear_model import LinearRegression
from sklearn.metrics import mean_absolute_error, mean_squared_error, r2_score
import pandas as pd # Optional: for loading data if in CSV
# Step 2: Load or generate data
# Example using generated data for simplicity
np.random.seed(42)
X = 2 * np.random.rand(100, 1) # 单个特征
y = 4 + 3 * X + np.random.randn(100, 1) # y = 4 + 3x + 噪声
# If loading from a file (like the original example):
# input_file = 'your_data_file.csv' # 替换为你的文件路径
# data = pd.read_csv(input_file)
# X = data[['feature_column']].values # 确保 X 是二维数组
# y = data['target_column'].values
# Step 3: Split data into training and testing sets
# test_size=0.3 means 30% for testing, 70% for training
# random_state ensures reproducibility
X_train, X_test, y_train, y_test = train_test_split(X, y, test_size=0.3, random_state=42)
# Step 4: Choose and instantiate the model
reg_linear = LinearRegression()
# Step 5: Train the model
reg_linear.fit(X_train, y_train)
# Step 6: Make predictions
y_pred = reg_linear.predict(X_test)
# Step 7: Evaluate the model
mae = mean_absolute_error(y_test, y_pred) # 平均绝对误差
mse = mean_squared_error(y_test, y_pred) # 均方误差
rmse = np.sqrt(mse) # 均方根误差
r2 = r2_score(y_test, y_pred) # R方
print("Regressor model performance:")
print(f"Mean Absolute Error (MAE): {mae:.4f}")
print(f"Mean Squared Error (MSE): {mse:.4f}")
print(f"Root Mean Squared Error (RMSE): {rmse:.4f}")
print(f"R-squared (R²): {r2:.4f}")
# Optional: Print learned parameters
print(f"\nIntercept (b₀): {reg_linear.intercept_[0]:.4f}") # 截距
print(f"Coefficient (b₁): {reg_linear.coef_[0][0]:.4f}") # 系数
# Step 8: Visualize the results
plt.figure(figsize=(8, 6))
sns.scatterplot(x=X_test.flatten(), y=y_test.flatten(), label='Actual Data') # 实际数据点
plt.plot(X_test, y_pred, color='red', linewidth=2, label='Regression Line') # 回归线
plt.title('Linear Regression Fit')
plt.xlabel('Feature (X)')
plt.ylabel('Target (y)')
plt.legend()
plt.show()

输出将显示:

  • 性能指标(MAE、MSE、RMSE、R²):较低的误差值(MAE、MSE、RMSE)和接近 1 的 R² 表示更好的性能。
  • 学习到的参数:线性模型的截距(Intercept)和系数(Coefficient(s))。
  • 显示实际数据点(散点图 Scatter plot)和拟合的回归线(Fitted regression line)的图。

(注意:原始教程的输出显示了一个负的 R²,表明对于该特定数据/拆分,拟合效果非常差。这强调了评估模型的重要性;简单的线性模型并非总是适用。)

除了简单线性回归和多元线性回归,其他重要算法包括:

  • 多项式回归(Polynomial Regression): 通过添加多项式项(如 x²、x³)来建模非线性关系。
  • Ridge 回归: 带有 L2 正则化(Regularization)的线性回归,通过惩罚大的系数来防止过拟合(Overfitting)。
  • Lasso 回归: 带有 L1 正则化的线性回归,还可以通过将一些系数收缩到零来实现特征选择(Feature selection)。
  • ElasticNet 回归: 结合了 L1 和 L2 正则化。
  • 支持向量回归(Support Vector Regression - SVR): 支持向量机(Support Vector Machines)在回归任务上的应用。
  • 决策树回归(Decision Tree Regression): 使用树状结构进行预测。
  • 随机森林回归(Random Forest Regression): 一种集成方法,使用多个决策树提高鲁棒性和准确性。
  • 梯度提升回归(Gradient Boosting Regression,例如 XGBoost、LightGBM、CatBoost): 强大的集成方法,顺序构建树,每棵树纠正前一棵树的错误。

回归在各种领域有着广泛的应用:

  • 预测分析/预测(Predictive Analysis/Forecasting): 预测未来值,如销量、股票价格、天气状况、能源需求或患者康复时间。
  • 理解关系(Understanding Relationships): 量化自变量对因变量的影响(例如,广告支出如何影响销售额,学习时间如何影响考试成绩)。
  • 优化(Optimization): 寻找流程的最佳设置(例如,基于需求预测优化定价)。
  • 风险评估(Risk Assessment): 预测金融风险、保险费或贷款违约概率。
  • 科学建模(Scientific Modeling): 建模物理现象、生物过程或经济趋势。

更多资源: