分组年度回归高效运行方案问询:大数据集内存问题解决
分组回归的内存与性能优化
问题背景
已实现按Year和Group分组执行简单线性回归,代码如下:
import pandas as pd import statsmodels.api as sm data = pd.DataFrame({ "Year": [2000, 2000, 2000, 2000, 2001, 2001, 2001, 2001], "Group": ["A", "A", "B", "B", "A", "A", "B", "B"], "ID": [1, 2, 3, 5, 2, 1, 2, 3], "Value 1": [40, 20, 30, 45, 22, 34, 11, 88], "Value 2": [3, 22, 11, 55, 5, 9, 4, 15], }) def func(group, var): X = group[var] # independent variable y = group["Value 2"] # dependent variable X = sm.add_constant(X) group["Residual"] = sm.OLS(y, X, missing="drop").fit().resid return group data.groupby(["Year", "Group"], group_keys=False).apply(func, var="Value 1")
在小数据集上运行正常,但在真实大数据集下运行缓慢,且出现内存错误:
MemoryError: Unable to allocate 5.73 GiB for an array with shape (229, 3359769) and data type float64
优化方案
1. 手动计算残差(推荐,内存占用极低)
对于仅含单个自变量的简单线性回归,无需调用statsmodels的完整OLS流程,直接利用统计量公式计算残差,避免构造大矩阵和不必要的中间变量:
import pandas as pd import numpy as np # 定义聚合函数,计算每个分组的关键统计量 def get_reg_stats(group): x = group["Value 1"] y = group["Value 2"] # 过滤缺失值 mask = ~x.isna() & ~y.isna() x_clean = x[mask] y_clean = y[mask] n = len(x_clean) if n < 2: # 样本量不足无法回归,残差设为NaN return pd.Series({"beta1": np.nan, "beta0": np.nan, "count": n}) mean_x = x_clean.mean() mean_y = y_clean.mean() cov_xy = (x_clean - mean_x).dot(y_clean - mean_y) / (n - 1) var_x = x_clean.var(ddof=1) beta1 = cov_xy / var_x beta0 = mean_y - beta1 * mean_x return pd.Series({"beta1": beta1, "beta0": beta0, "count": n}) # 按分组计算回归参数 reg_params = data.groupby(["Year", "Group"]).apply(get_reg_stats).reset_index() # 合并参数到原数据,计算残差 data = data.merge(reg_params, on=["Year", "Group"], how="left") data["Residual"] = data["Value 2"] - (data["beta0"] + data["beta1"] * data["Value 1"]) # 清理临时列 data = data.drop(columns=["beta0", "beta1", "count"])
优势:
- 仅计算必要的统计量(均值、协方差、方差),内存占用仅为原数据量级,不会生成大矩阵
- 避免
statsmodels拟合过程中的额外内存开销,速度提升显著 - 自动处理样本量不足的分组,残差设为NaN
2. 分块处理大数据
若数据集过大无法一次性载入内存,可使用pandas的分块读取功能,逐块处理后合并结果:
chunk_size = 100000 # 根据内存调整块大小 chunks = [] for chunk in pd.read_csv("your_large_data.csv", chunksize=chunk_size): # 对当前块执行上述手动残差计算逻辑 reg_params = chunk.groupby(["Year", "Group"]).apply(get_reg_stats).reset_index() chunk = chunk.merge(reg_params, on=["Year", "Group"], how="left") chunk["Residual"] = chunk["Value 2"] - (chunk["beta0"] + chunk["beta1"] * chunk["Value 1"]) chunk = chunk.drop(columns=["beta0", "beta1", "count"]) chunks.append(chunk) # 合并所有块 final_data = pd.concat(chunks, ignore_index=True)
3. 优化原statsmodels调用(仅适用于稍大数据,不推荐超大规模)
若坚持使用statsmodels,可避免在分组内修改原数据,直接返回残差序列,减少内存复制:
def func(group, var): x = group[var].dropna() y = group["Value 2"].dropna() if len(x) < 2: return pd.Series(np.nan, index=group.index, name="Residual") X = sm.add_constant(x) resid = sm.OLS(y, X).fit().resid # 对齐索引,确保残差与原分组行匹配 resid_full = pd.Series(np.nan, index=group.index, name="Residual") resid_full.loc[resid.index] = resid return resid_full # 直接生成残差列,避免整个分组的复制 data["Residual"] = data.groupby(["Year", "Group"], group_keys=False).apply(func, var="Value 1")
注意:此方法仍会生成设计矩阵,内存占用较高,仅适用于分组规模较小的场景。
内容的提问来源于stack exchange,提问作者Derek
相关产品推荐
相关产品推荐

