You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

分组年度回归高效运行方案问询:大数据集内存问题解决

分组回归的内存与性能优化

问题背景

已实现按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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.19 11:37:27