PyTorch中如何避免重复计算实现梯度依赖部分剥离?
问题:避免重复计算带detach的对数似然值
我正在构建一个概率模型$q_{\phi}(x)$,它能采样样本点并返回对应的似然值。下面是一个高斯分布的示例实现:
import torch import torch.nn as nn import math as m class NormalModel(nn.Module): def __init__(self, mu : torch.Tensor, sigma : torch.Tensor): super().__init__() self.dim_ = mu.shape[0] self.mu_ = nn.Parameter(mu, requires_grad=True) self.sigma_ = nn.Parameter(sigma, requires_grad=True) def sample(self, n_sample): eps = torch.normal(mean=torch.zeros(n_sample, self.dim_), std=torch.ones(n_sample, self.dim_)) sample = self.mu_ + eps * self.sigma_ # 普通对数似然(样本依赖参数,梯度会传播到mu和sigma) log_prob = -(0.5 * ((sample - self.mu_)/self.sigma_)**2 + torch.log(self.sigma_) + m.log(2*m.pi)).sum(dim=-1) # 分离样本后的对数似然(样本不参与梯度传播,仅mu和sigma的梯度保留) log_prob_detach = -(0.5 * ((sample.detach() - self.mu_)/self.sigma_)**2 + torch.log(self.sigma_) + m.log(2*m.pi)).sum(dim=-1) return sample, log_prob, log_prob_detach
我的需求是构建同时包含两种似然的损失函数:一种是样本依赖模型参数的$q_{\phi}(x_{\phi})$,另一种是样本不依赖参数的$q_{\phi}(x)$。目前的实现需要重复计算对数似然的表达式,只在sample是否detach上做区分。有没有办法只计算一次核心表达式,再通过某种操作得到log_prob_detach,从而避免重复计算?
解决方案:拆分计算项,选择性截断梯度
当然可以,通过拆分对数似然的计算逻辑,把依赖样本和仅依赖模型参数的部分分开,就能避免重复计算共享项,同时实现梯度的选择性截断:
import torch import torch.nn as nn import math as m class NormalModel(nn.Module): def __init__(self, mu : torch.Tensor, sigma : torch.Tensor): super().__init__() self.dim_ = mu.shape[0] self.mu_ = nn.Parameter(mu, requires_grad=True) self.sigma_ = nn.Parameter(sigma, requires_grad=True) # 预计算常数项,避免重复计算 self.const_term = m.log(2 * m.pi) def sample(self, n_sample): eps = torch.normal(mean=torch.zeros(n_sample, self.dim_), std=torch.ones(n_sample, self.dim_)) sample = self.mu_ + eps * self.sigma_ # 拆分计算:分离样本相关项和参数相关项 centered = sample - self.mu_ scaled = centered / self.sigma_ sample_dependent = 0.5 * scaled ** 2 param_dependent = torch.log(self.sigma_) + self.const_term # 常规对数似然:保留所有梯度 log_prob = -(sample_dependent + param_dependent).sum(dim=-1) # 构造detach版本:仅截断样本相关部分的梯度,参数部分保留梯度 sample_dependent_detached = 0.5 * (centered.detach() / self.sigma_) ** 2 log_prob_detach = -(sample_dependent_detached + param_dependent).sum(dim=-1) return sample, log_prob, log_prob_detach
关键思路
- 拆分计算逻辑:把对数似然拆成两部分:和样本直接相关的
sample_dependent,以及只和模型参数(mu、sigma)相关的param_dependent。后者只需要计算一次,就能同时用于两种log_prob的计算。 - 选择性detach:只对样本相关的张量执行
detach(),这样log_prob_detach就不会从样本反向传播梯度,但依然保留对模型参数的梯度计算能力,完全符合需求。
如果追求更简洁的写法,也可以基于已计算的log_prob来调整梯度,但拆分计算项的方式更直观,也更容易调试和维护。
内容的提问来源于stack exchange,提问作者RobVerheyen
相关产品推荐
相关产品推荐

