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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.22 12:15:42