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

PyTorch使用autograd时如何为自定义计算步骤手动添加梯度

解决方案:自定义Autograd算子实现手动梯度注入

PyTorch原生支持自定义可微算子,完全兼容现有autograd自动求导机制,不需要全链路手动写梯度计算逻辑,仅需要为无法被追踪的中间步骤手动补充forward和backward实现即可。

实现逻辑

  • 继承torch.autograd.Function类实现自定义算子:
    • forward方法中写入你无法被autograd追踪的计算逻辑,输入为需要参与梯度回传的所有张量(本场景为m、s、p),计算时可以正常detach转成numpy调用scipy方法,最终返回计算结果(本场景为归一化后的p),同时可以把forward阶段需要用到的中间变量存到ctx上下文对象中,供backward阶段使用
    • backward方法中写入你手动推导的梯度公式,输入为上游回传的损失对算子输出的梯度,计算出损失对算子输入的所有张量的梯度后返回即可,返回值顺序需要和forward的输入参数顺序一一对应
  • 后续使用时直接调用自定义算子的apply方法传入参数即可,算子会自动嵌入autograd计算图,调用loss.backward()时会自动调用你写的backward逻辑完成梯度传播

本场景的代码实现

首先确认梯度推导逻辑:定义每个高斯分量的区间积分值为 cdf_diff_i = Φ(higherCutoff, m_i, s_i) - Φ(lowerCutoff, m_i, s_i),其中Φ为正态分布CDF,φ为正态分布PDF;归一化系数 normFactor = sum(p_i * cdf_diff_i),归一化后权重 p_norm_i = p_i / normFactor

import torch
import scipy.stats as spstats

class NormalizedWeight(torch.autograd.Function):
    @staticmethod
    def forward(ctx, m, s, p, higher_cutoff, lower_cutoff, eps=1e-6):
        # 输入m、s、p的形状均为 [1, numberGaussians]
        device = m.device
        m_np = m[0].detach().cpu().numpy()
        s_np = s[0].detach().cpu().numpy() + eps
        p_np = p[0].detach().cpu().numpy()
        
        # 计算每个分量的CDF差值
        cdf_high = spstats.norm.cdf(higher_cutoff, m_np, s_np)
        cdf_low = spstats.norm.cdf(lower_cutoff, m_np, s_np)
        cdf_diff = cdf_high - cdf_low
        # 计算归一化系数
        norm_factor = (p_np * cdf_diff).sum()
        # 计算归一化后的p
        p_norm = p / norm_factor
        
        # 保存backward需要的中间变量到ctx
        ctx.save_for_backward(m, s, p)
        ctx.aux_params = (higher_cutoff, lower_cutoff, eps, cdf_diff, norm_factor)
        return p_norm
    
    @staticmethod
    def backward(ctx, grad_output):
        # 取出保存的变量
        m, s, p = ctx.saved_tensors
        higher_cutoff, lower_cutoff, eps, cdf_diff, norm_factor = ctx.aux_params
        device = m.device
        
        # 转numpy计算PDF相关值
        m_np = m[0].detach().cpu().numpy()
        s_np = s[0].detach().cpu().numpy() + eps
        p_np = p[0].detach().cpu().numpy()
        
        # 计算需要的PDF值
        pdf_high = spstats.norm.pdf(higher_cutoff, m_np, s_np)
        pdf_low = spstats.norm.pdf(lower_cutoff, m_np, s_np)
        # 转成torch张量方便计算
        cdf_diff_t = torch.tensor(cdf_diff, device=device)[None, :]
        pdf_high_t = torch.tensor(pdf_high, device=device)[None, :]
        pdf_low_t = torch.tensor(pdf_low, device=device)[None, :]
        
        # 计算对原始p的梯度
        grad_p = grad_output * (1/norm_factor - p * cdf_diff_t / (norm_factor **2))
        
        # 计算对m的梯度:先算normFactor对m的偏导
        d_norm_dm = p * (pdf_low_t - pdf_high_t) / s
        grad_m = grad_output * (- p / (norm_factor **2)) * d_norm_dm
        
        # 计算对s的梯度:先算normFactor对s的偏导
        term_high = (higher_cutoff - m) * pdf_high_t / (s **2)
        term_low = (lower_cutoff - m) * pdf_low_t / (s **2)
        d_norm_ds = p * (term_low - term_high)
        grad_s = grad_output * (- p / (norm_factor **2)) * d_norm_ds
        
        # 后面的截断值、eps等固定参数不需要梯度,返回None占位
        return grad_m, grad_s, grad_p, None, None, None

调用方法

替换你原来的归一化代码段即可,后续正常计算损失调用backward就能得到正确的梯度:

# 假设m、s、p是你模型输出的张量,requires_grad=True
p_norm = NormalizedWeight.apply(m, s, p, higherCutoff, lowerCutoff)
# 后续用p_norm计算损失,调用loss.backward()即可

内容的提问来源于stack exchange,提问作者Sam

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.03 15:48:01