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

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:存储每个时间步的η_k
  • self.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

优化说明

  1. 避免递归:不再用嵌套lambda,通过存储累积参数直接计算结果,彻底解决递归深度问题。
  2. 高效更新:每次update仅需O(1)时间的乘法和追加操作,无需遍历历史数据。
  3. 批量友好:基于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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.01 19:31:22