Python迭代计算最小二乘系数:如何进一步提升代码运行速度?
问题:OLS系数计算的性能优化空间
每个时间步t会获取新的x_t(对应金融术语中的benchmark)和y_t(对应asset),当累计的x、y数据总量达到PERIOD时,执行普通最小二乘法(OLS)并输出斜率系数。
此前使用scipy.stats.linregress获取系数,但速度比当前实现慢约50%——推测是因为该方法额外计算了截距、p值等指标。现在希望挖掘当前代码的进一步提速空间。
现有代码实现
import numpy as np class Beta(): def __init__(self, PERIOD): self.PERIOD = PERIOD self.stat = None self.benchmark_vec = [] self.asset_vec = [] def update(self, new_benchmark, new_asset): """ assumes that on update, neither new_asset nor new_benchmark are none """ if new_benchmark is not None and new_asset is not None: self.benchmark_vec.append(new_benchmark) self.asset_vec.append(new_asset) if len(self.benchmark_vec)>self.PERIOD: self.benchmark_vec=self.benchmark_vec[1:] if len(self.asset_vec)>self.PERIOD: self.asset_vec=self.asset_vec[1:] if len(self.benchmark_vec)==self.PERIOD and len(self.asset_vec)==self.PERIOD: self.stat=self.online_calc() return self.stat def online_calc(self): asset_ar = np.array(self.asset_vec) benchmark_ar = np.array(self.benchmark_vec) y = asset_ar[1:]/asset_ar[:-1]-1 x = benchmark_ar[1:]/benchmark_ar[:-1]-1 X = x.reshape(-1, 1) X = np.concatenate([np.ones_like(X), X], axis=1) b = np.linalg.pinv(X.T.dot(X)).dot(X.T).dot(y) return b[1] ###### example usage asset_vec=[ 0.01568653, -0.00669479, 0.01140213, -0.00107317, -0.00131155, -0.00333463, -0.00114006, 0.00263075, -0.00507337, 0.00712401, 0.00388323] benchmark_vec=[ 0.01150227, 0.00045742, 0.01114376, -0.00305367, 0.00388245, 0.00323491, -0.00449446, -0.00075698, -0.01114904, 0.0147878, 0.00528754] beta = Beta(4) output = [] for b, a in zip(benchmark_vec, asset_vec): output.append(beta.update(b, a)) print(output)
优化方向及实现
1. 核心性能损耗点
当前代码的主要低效之处:
- 每次计算都将列表转换为numpy数组,重复开销大
- 重新计算整个窗口的收益率,而非增量更新
- 使用矩阵伪逆求解OLS,存在运算冗余(仅需斜率系数)
- 列表切片维护窗口,大
PERIOD下效率低
2. 优化后的代码
import numpy as np from collections import deque class OptimizedBeta(): def __init__(self, PERIOD): self.PERIOD = PERIOD self.stat = None # 维护原始价格的固定长度窗口 self.benchmark_prices = deque(maxlen=PERIOD) self.asset_prices = deque(maxlen=PERIOD) # 维护收益率的固定长度窗口(长度为PERIOD-1) self.x_returns = deque(maxlen=PERIOD-1) self.y_returns = deque(maxlen=PERIOD-1) # 预计算统计量,避免重复遍历 self._sum_x = 0.0 self._sum_y = 0.0 self._sum_x2 = 0.0 self._sum_xy = 0.0 def update(self, new_benchmark, new_asset): if new_benchmark is None or new_asset is None: return self.stat prev_benchmark = self.benchmark_prices[-1] if len(self.benchmark_prices) > 0 else None prev_asset = self.asset_prices[-1] if len(self.asset_prices) > 0 else None # 添加新价格到窗口 self.benchmark_prices.append(new_benchmark) self.asset_prices.append(new_asset) # 计算新收益率(当有前一个价格时) if prev_benchmark is not None and prev_asset is not None: x = (new_benchmark / prev_benchmark) - 1 y = (new_asset / prev_asset) - 1 # 窗口已满时,先移除最旧的收益率并更新统计量 if len(self.x_returns) == self.PERIOD - 1: old_x = self.x_returns.popleft() old_y = self.y_returns.popleft() self._sum_x -= old_x self._sum_y -= old_y self._sum_x2 -= old_x ** 2 self._sum_xy -= old_x * old_y # 添加新收益率并更新统计量 self.x_returns.append(x) self.y_returns.append(y) self._sum_x += x self._sum_y += y self._sum_x2 += x ** 2 self._sum_xy += x * y # 收益率窗口达标时计算系数 if len(self.x_returns) == self.PERIOD - 1: n = len(self.x_returns) # 直接用斜率公式:beta = Cov(x,y)/Var(x) cov_xy = (self._sum_xy / n) - (self._sum_x / n) * (self._sum_y / n) var_x = (self._sum_x2 / n) - (self._sum_x / n) ** 2 # 避免除以0的边界情况 self.stat = cov_xy / var_x if var_x != 0 else 0.0 else: self.stat = None return self.stat # 示例使用 if __name__ == "__main__": asset_vec=[ 0.01568653, -0.00669479, 0.01140213, -0.00107317, -0.00131155, -0.00333463, -0.00114006, 0.00263075, -0.00507337, 0.00712401, 0.00388323] benchmark_vec=[ 0.01150227, 0.00045742, 0.01114376, -0.00305367, 0.00388245, 0.00323491, -0.00449446, -0.00075698, -0.01114904, 0.0147878, 0.00528754] beta = OptimizedBeta(4) output = [] for b, a in zip(benchmark_vec, asset_vec): output.append(beta.update(b, a)) print(output)
3. 优化点说明
- 用deque维护固定长度窗口:避免列表切片的O(n)开销,弹出头部元素为O(1)
- 增量计算收益率与统计量:每次仅计算新增的收益率,维护总和、平方和、交叉乘积和,避免重复遍历数组
- 直接使用斜率公式:跳过矩阵构造与伪逆求解,利用协方差/方差的公式直接计算斜率,运算量大幅降低
- 减少数组转换:全程使用deque和基础数值类型维护数据,避免频繁的列表转numpy数组操作
内容的提问来源于stack exchange,提问作者cmaz
相关产品推荐
相关产品推荐

