PyTorch中高效定义随时间步更新的递归概率函数求助
PyTorch中高效实现随时间步更新的递归概率函数
问题分析
你遇到的RecursionError是因为每次调用update时,都把self.log_y赋值为嵌套lambda,这个lambda会递归调用之前版本的self.log_y,随着更新次数增加,递归深度不断累积,最终超出Python的递归深度限制。
另外,直接遍历存储参数的列表计算会导致重复计算,效率低下。我们可以通过展开对数形式的递推公式,用累积量存储中间结果,实现O(1)时间的更新和高效的批量计算。
递推公式展开(对数形式)
原始递推的对数形式为:
log_y_{t+1}(x) = [log_y_t(x) - η_t * f_t(x)] / (1 + η_t * λ_t)
展开多步递推后可以得到:
log_y_T(x) = log_y_0(x) / (Π_{k=0}^{T-1} C_k) - Σ_{k=0}^{T-1} [ η_k * f_k(x) / (Π_{m=k}^{T-1} C_m) ]
其中C_k = 1 + η_k * λ_k,Π表示乘积运算。
高效实现方案
我们可以维护几个核心变量,通过累积量避免递归和重复计算:
self.log_y0:初始的对数概率函数self.prod_C:所有C_k的累积乘积(Π_{k=0}^{t-1} C_k)self.etas:存储每个时间步的η_kself.prod_C_up_to:存储到每个时间步k之前的累积乘积(Π_{m=0}^{k-1} C_m),初始为[1]
每次更新仅需简单的乘法和追加操作,计算对数概率时直接利用这些累积量组合结果。
代码实现
import torch class RecursiveProbability: def __init__(self, log_y0): self.log_y0 = log_y0 # 初始对数概率函数,输入x返回对应log_y0(x) self.prod_C = 1.0 # 累积乘积,初始为1(对应0步更新) self.etas = [] # 存储各时间步的eta_t self.prod_C_up_to = [1.0] # 存储到每个时间步之前的累积乘积,初始为[1] def update(self, eta_t, lambda_t): # 计算当前时间步的C_t C_t = 1 + eta_t * lambda_t # 更新累积乘积 self.prod_C *= C_t # 存储当前eta_t self.etas.append(eta_t) # 追加到prod_C_up_to(到当前步的累积乘积) self.prod_C_up_to.append(self.prod_C) def compute_log_y(self, x, f_list): """ 计算当前时间步的log_y(x) :param x: 输入张量 :param f_list: 存储各时间步f_t(x)的函数列表,顺序与update的时间步一致 """ # 计算初始项:log_y0(x) / 累积乘积 initial_term = self.log_y0(x) / self.prod_C # 计算求和项:Σ(eta_k * f_k(x) * (prod_C_up_to[k]/prod_C)) sum_terms = torch.zeros_like(initial_term) for k in range(len(self.etas)): weight = self.prod_C_up_to[k] / self.prod_C sum_terms += self.etas[k] * weight * f_list[k](x) # 最终log_y(x) return initial_term - sum_terms
优化说明
- 避免递归:不再用嵌套lambda,通过存储累积参数直接计算结果,彻底解决递归深度问题。
- 高效更新:每次
update仅需O(1)时间的乘法和追加操作,无需遍历历史数据。 - 批量友好:基于PyTorch张量运算,支持批量输入x,可利用GPU加速,远快于遍历列表的逐元素计算。
使用示例
# 定义初始对数概率函数 def initial_log_y(x): return torch.log(torch.sigmoid(x)) # 示例:sigmoid的对数 # 初始化实例 prob_model = RecursiveProbability(initial_log_y) # 模拟三次更新 f0 = lambda x: x**2 f1 = lambda x: torch.sin(x) f2 = lambda x: torch.exp(-x) prob_model.update(eta_t=0.1, lambda_t=0.05) prob_model.update(eta_t=0.2, lambda_t=0.1) prob_model.update(eta_t=0.15, lambda_t=0.08) # 计算log_y(x) x = torch.randn(10, 5) # 批量输入 log_y_result = prob_model.compute_log_y(x, [f0, f1, f2]) print(log_y_result.shape) # 输出torch.Size([10,5])
内容的提问来源于stack exchange,提问作者Vimar
相关产品推荐
相关产品推荐

