Pyro训练可训练伯努利分布时触发反向传播二次计算RuntimeError
问题解决方案
报错核心原因
你遇到的RuntimeError是计算图缓存释放导致的:
- 伯努利分布仅在训练循环外初始化了一次,第一次迭代调用
loss.backward()后,分布计算log_prob用到的中间计算缓存会被自动释放,第二次迭代反向传播时会尝试访问已经被释放的缓存节点,触发报错。 - 除此之外你的训练逻辑还缺少梯度清零、参数更新的步骤,无法完成正常的参数迭代。
修复步骤
- 将伯努利分布的初始化挪到训练循环内部,每次迭代用最新的参数创建新的分布对象,保证计算图每次都是全新的
- 新增优化器梯度清零
optim.zero_grad()和参数更新optim.step()逻辑,完善训练流程 - 提前将稀疏的训练数据转为稠密PyTorch张量,避免每次迭代重复转换带来的性能损耗
- 先设置参数的
requires_grad=True再传入优化器,保证优化器可以正确跟踪可训练参数
修正后完整代码
import torch import pyro pyd = pyro.distributions print("torch version:", torch.__version__) print("pyro version:", pyro.__version__) import numpy as np torch.manual_seed(123) # 1. 初始化可训练参数 train_vars = (pyd.Uniform(low=torch.FloatTensor([0.01]), high=torch.FloatTensor([0.1])).rsample([train_data.shape[-1]]).squeeze()) # 先设置参数可训练 train_vars.requires_grad = True # 2. 预处理类别0的训练数据:稀疏矩阵先转稠密再转张量 class_mask = (train_labels == 0) # 若train_data是scipy稀疏矩阵,用.todense()转为稠密数组 class_data = torch.tensor(train_data[class_mask, :].todense(), dtype=torch.float) # 3. 初始化优化器 optim = torch.optim.Adam([train_vars]) # 4. 训练循环 for i in range(100): # 每次迭代用最新参数创建伯努利分布 distribution = pyd.Bernoulli(probs=train_vars) # 计算NLL损失 loss = -torch.mean(distribution.log_prob(class_data)) # 梯度清零 optim.zero_grad() # 反向传播 loss.backward() # 更新参数 optim.step() # 可选:打印训练进度 if i % 10 == 0: print(f"迭代次数 {i:03d} | NLL损失: {loss.item():.4f}")
注意:不要直接给
backward()添加retain_graph=True参数,这会导致计算图不断积累占用内存,属于治标不治本的方案。
内容的提问来源于stack exchange,提问作者js kim
相关产品推荐
相关产品推荐

