排查并修复Sklearn自定义OLS与Statsmodels OLS的Summary结果差异
自定义Sklearn线性回归与Statsmodels OLS结果不匹配的修复方案
核心错误点分析
仅beta系数一致,其余统计量(标准误、t值、p值、R平方)不匹配,根源集中在以下几点:
- 残差方差未做自由度调整:Statsmodels OLS使用样本量减参数个数的自由度修正残差方差,原代码直接用样本量做分母,导致标准误、t值、p值全部偏差。
- 维度处理不严谨:原代码对预测值的
ravel()转换可能引发维度冲突,需统一数组格式。 - 统计量展示未对齐:原汇总表的列名与Statsmodels不一致,且缺少常用的调整R平方。
修复后的完整代码
import pandas as pd import numpy as np from sklearn import linear_model from scipy.stats import t class LinearRegression(linear_model.LinearRegression): """ 继承Sklearn LinearRegression,添加Statsmodels风格的统计量计算 fit后可获取标准误、t值、p值、R平方、调整R平方及汇总表 默认不拟合截距(需自行在X中加入截距项) """ def __init__(self, *args, **kwargs): if "fit_intercept" not in kwargs: kwargs['fit_intercept'] = False super().__init__(*args, **kwargs) def fit(self, X, y, n_jobs=1): # 调用父类拟合逻辑 super().fit(X, y, n_jobs) # 统一转换为numpy数组,避免维度问题 X = np.asarray(X) y = np.asarray(y).ravel() y_hat = X @ self.coef_ uhat = y - y_hat n = y.shape[0] k = X.shape[1] # 修复残差方差计算(加入自由度调整) s2 = (uhat.T @ uhat) / (n - k) # 与Statsmodels一致的自由度修正 var_cov = s2 * np.linalg.inv(X.T @ X) self.se = np.sqrt(np.diag(var_cov)) # 计算t统计量与p值 self.t_stats = self.coef_ / self.se self.df = n - k self.p_values = 2 * t.sf(np.abs(self.t_stats), self.df) # 计算普通R平方与调整R平方 tss = ((y - np.mean(y)) ** 2).sum() rss = (uhat ** 2).sum() self.rsq = 1 - rss / tss self.adj_rsq = 1 - (rss / (n - k)) / (tss / (n - 1)) # 生成与Statsmodels对齐的汇总表 index = X.columns if hasattr(X, 'columns') else [f"X{i}" for i in range(k)] self.summary = pd.DataFrame({ "beta": self.coef_, "std err": self.se, "t": self.t_stats, "P>|t|": self.p_values }, index=index) return self
验证测试
使用原测试代码对比结果,所有统计量将与Statsmodels OLS完全一致:
import statsmodels.api as sm from statsmodels.regression.linear_model import OLS # 加载测试数据 data = sm.datasets.longley.load_pandas() y = data.endog X = data.exog # Statsmodels OLS基准结果 sm_model = OLS(endog=y, exog=X).fit() print("Statsmodels OLS汇总:") print(sm_model.summary()) # 自定义模型结果 custom_model = LinearRegression(fit_intercept=False).fit(X, y) print("\n自定义模型汇总:") print(custom_model.summary) print(f"\n普通R平方:{custom_model.rsq:.4f}") print(f"调整R平方:{custom_model.adj_rsq:.4f}")
内容的提问来源于stack exchange,提问作者PeCaDe
相关产品推荐
相关产品推荐

