Python datatable按user_id分组应用自定义OLS函数计算评分趋势方案咨询
解决方案
前提注意事项
你当前代码中使用range(s)作为OLS的自变量,本质是默认每个用户的评分记录按时间先后排序,因此首先必须对全表按user_id、date排序,保证每个用户分组内的rating序列和时间顺序一致,否则计算出的趋势斜率无意义。
datatable原生实现方案
你之前的写法报错的核心原因是:datatable默认j位置的表达式是逐行计算,直接传入dt.f.rating会让coef()函数接收到单个数值而非整个分组的序列,需要用dt.apply()标记为组级聚合函数,完整代码如下:
import datatable as dt import numpy as np # 1. 按用户、日期排序 mydt = mydt.sort(dt.f.user_id, dt.f.date) # 2. 调整自定义函数,增加边界判断避免样本不足报错 def calculate_trend(y): s = len(y) if s < 2: return np.nan A = np.vstack([np.arange(s), np.ones(s)]).T m, _ = np.linalg.lstsq(A, y, rcond=None)[0] return m # 3. 分组聚合计算斜率 result = mydt[:, dt.apply(calculate_trend, dt.f.rating), by(dt.f.user_id)]
计算完成后result表会包含两列:user_id和对应计算出的趋势斜率。
性能优化方案
如果数据量极大,Python层自定义函数的开销会比较高,可以用Numba编译函数,同时手写OLS斜率公式替代np.linalg.lstsq,性能可以提升数倍到数十倍:
from numba import jit @jit(nopython=True) def calculate_trend_numba(y): s = len(y) if s < 2: return np.nan x = np.arange(s) x_mean = x.mean() y_mean = y.mean() cov = ((x - x_mean) * (y - y_mean)).sum() var = ((x - x_mean) ** 2).sum() return cov / var if var != 0 else np.nan # 调用方式不变 result = mydt[:, dt.apply(calculate_trend_numba, dt.f.rating), by(dt.f.user_id)]
替代方案
如果datatable的分组聚合性能仍不符合预期,可以换用Polars库,其对分组自定义函数的优化程度更高,且同样支持内存友好的大数据处理,参考代码如下:
import polars as pl # 直接读入原始数据,支持懒加载无需全量加载进内存 mydf = pl.scan_csv("你的数据文件路径") result = ( mydf .sort(["user_id", "date"]) .group_by("user_id") .agg( pl.col("rating").map_elements(calculate_trend_numba, return_dtype=pl.Float64).alias("trend_coef") ) .collect() # 最后才执行计算,内存占用极低 )
内容的提问来源于stack exchange,提问作者Areza
相关产品推荐
相关产品推荐

