在scikit-learn中实现数据子集独有斜率截距的复合线性回归模型
模型归类说明
你要实现的这个模型属于**分层线性模型(Hierarchical Linear Model, HLM,也叫多水平线性模型)**的变种,核心特征是同时拟合全局共享的特征系数P_k,以及每个分组(这里是站点)独有的偏移S_0n和缩放系数S_1n,既能保留全局特征的共性规律,又能适配不同站点的分布差异。
实现方案
方案1:自定义scikit-learn评估器(推荐,适配你后续统一测试其他算法的需求)
你不需要用Pipeline,直接继承sklearn.base.BaseEstimator和RegressorMixin即可,完全符合scikit-learn的接口规范,后续可以直接和GridSearchCV、cross_val_score等工具搭配使用,步骤如下:
- 第一步:参数初始化时指定站点数量、特征数量
- 第二步:fit方法内部用
scipy.optimize.minimize求解损失函数,损失函数默认用均方误差即可,有正则需求可以额外加L1/L2惩罚项,待优化参数包含全局P_k(k+1个,含截距P0)、每个站点的S0n和S1n(共2n个) - 第三步:predict方法里先读取输入样本所属站点,带入对应S0n、S1n和全局P计算预测值
简化代码模板参考:
from sklearn.base import BaseEstimator, RegressorMixin import numpy as np from scipy.optimize import minimize class SiteAdaptedLinearRegressor(BaseEstimator, RegressorMixin): def __init__(self, n_sites, n_features): self.n_sites = n_sites self.n_features = n_features # 不含截距的原始特征数量 def fit(self, X, y): # 输入X约定:前n_features列是原始特征,最后1列是站点编号(取值为0到n_sites-1的整数) site_ids = X[:, -1].astype(int) X_feat = X[:, :-1] # 初始化参数:P初始为0,S0初始为0,S1初始为1(符合无调整的默认场景) init_params = np.concatenate([ np.zeros(self.n_features + 1), np.zeros(self.n_sites), np.ones(self.n_sites) ]) # 定义损失函数 def loss(params): P = params[:self.n_features+1] S0 = params[self.n_features+1 : self.n_features+1+self.n_sites] S1 = params[self.n_features+1+self.n_sites : ] # 计算全局线性输出 global_pred = P[0] + X_feat @ P[1:] # 按站点做缩放偏移调整 site_pred = S0[site_ids] + S1[site_ids] * global_pred return np.mean((site_pred - y)**2) # 可自行添加正则项,比如+ 0.01*np.sum(P**2) # 求解最优参数 res = minimize(loss, init_params, method='L-BFGS-B') self.P_ = res.x[:self.n_features+1] self.S0_ = res.x[self.n_features+1 : self.n_features+1+self.n_sites] self.S1_ = res.x[self.n_features+1+self.n_sites : ] return self def predict(self, X): site_ids = X[:, -1].astype(int) X_feat = X[:, :-1] global_pred = self.P_[0] + X_feat @ self.P_[1:] return self.S0_[site_ids] + self.S1_[site_ids] * global_pred
方案2:直接用scipy.optimize求解
如果你不需要和scikit-learn的其他工具配套使用,也可以直接单独编写上述损失函数和求解逻辑,核心逻辑和自定义评估器完全一致,仅省略了接口适配的部分。
替代方案:特征工程适配普通线性回归
如果你的站点数量很少,也可以通过特征构造直接用scikit-learn自带的LinearRegression实现:
- 对站点ID做独热编码,得到n个0-1特征
- 构造两类新特征:所有站点独热编码本身、每个原始特征和所有站点独热编码的乘积
- 直接用普通线性回归拟合即可,拟合后的参数可以拆解为你需要的S0n、S1n和全局P组合,优点是不需要自定义求解逻辑,缺点是站点数量多时特征维度会爆炸,拟合速度大幅下降。
内容的提问来源于stack exchange,提问作者RobinSheehy
相关产品推荐
相关产品推荐

