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
相关产品推荐
相关产品推荐

