PyTorch自定义激活函数:类实现与函数实现的差异探究
我在PyTorch里实现自定义激活函数时发现,官方文档建议用torch.autograd.Function类的方式(示例(b)),但用普通函数实现(示例(a))也没报错。想搞清楚这两种方式在反向传播上的差异,示例(a)的写法是否合适、反向传播能否正常进行?如果可以,为什么不会报错?
另外我实际需要实现一个无导数操作:把复数的振幅强制设为1(示例(c)),目前用示例(b)的思路,想知道示例(a)的方式能不能用,以及正确的写法是什么。
示例(a):普通函数实现自定义激活函数
import torch import torch.nn as nn def mytanh(x): xamp = torch.abs(x) xphs = torch.angle(x) xamp = torch.tanh(xamp) y = xamp*torch.exp(1j*xphs) return y class NetEn(nn.Module): def __init__(self): super(NetEn, self).__init__() self.conv1 = nn.Conv2d(1, 64, 5, 2, 0) def forward(self, x): x = self.conv1(x) x = mytanh(x) return x
示例(b):继承autograd.Function类实现自定义激活函数
class MyTanh(torch.autograd.Function): @staticmethod def forward(ctx, x): # context 状态 ctx.save_for_backward(x) xamp = torch.abs(x) xphs = torch.angle(x) xamp = torch.tanh(xamp) y = xamp*torch.exp(1j*xphs) return y @staticmethod def backward(ctx, dL_dy): # dL_dy = dL/dy(梯度输出) x, = ctx.saved_tensors dy_dx = 1 - torch.tanh(x)**2 dL_dx = dL_dy * dy_dx return dL_dx # 梯度输入 class NetEn(nn.Module): def __init__(self): super(NetEn, self).__init__() self.conv1 = nn.Conv2d(1, 64, 5, 2, 0) def forward(self, x): mytanh = MyTanh.apply x = self.conv1(x) x = mytanh(x) return x
示例(c):普通函数实现振幅设为1的操作
def abscut(xcomp): xcomp_o = torch.exp(1j*(torch.angle(xcomp))) return xcomp_o class NetEn(nn.Module): def __init__(self): super(NetEn, self).__init__() self.conv1 = nn.Conv2d(1, 64, 5, 2, 0) def forward(self, x): x = self.conv1(x) x = abscut(x) return x
核心解答
1. 示例(a)的普通函数写法为什么能正常反向传播?
你用普通函数实现时没报错,是因为函数内部调用的torch.abs、torch.angle、torch.tanh、torch.exp都是PyTorch提供的可微分操作——这些函数都自带了反向传播的梯度计算逻辑。PyTorch的autograd会自动追踪这些操作,构建计算图,在反向传播时自动计算梯度,所以完全能正常运行。
这种写法适合:你的自定义操作完全由PyTorch内置可微分函数组合而成,不需要手动修改梯度逻辑的场景。
2. 示例(a)和示例(b)的差异
- 示例(a):依赖PyTorch自动生成梯度,代码更简洁,不需要手动写backward逻辑,但无法自定义梯度计算规则。
- 示例(b):手动实现forward和backward,完全掌控梯度计算过程。适合两种场景:
- 内置函数的梯度不符合你的需求,需要自定义;
- 操作本身不可微分(比如示例(c)的情况),需要手动定义反向传播的行为(比如把梯度设为0,或者用近似梯度)。
注意:你示例(b)里的backward逻辑可能有问题——原mytanh操作是针对复数的振幅做tanh,你的backward计算的1 - torch.tanh(x)**2是实数tanh的导数,但复数输入的话,这个梯度计算是否正确需要验证。如果用示例(a)的写法,PyTorch会自动处理复数的梯度,更稳妥。
3. 示例(c)的无导数操作怎么写才正确?
示例(c)的普通函数写法有问题:虽然forward能得到振幅为1的复数,但反向传播时,torch.angle对复数输入的导数在PyTorch里是有定义的,autograd会自动计算梯度,这和你想要的“无导数操作”(即该操作不影响上游参数的梯度)不符。
正确的实现方式是用torch.autograd.Function自定义,手动把backward的梯度设为0,或者根据需求定义:
class AbsCut(torch.autograd.Function): @staticmethod def forward(ctx, x): # 把振幅设为1,保留相位 return torch.exp(1j * torch.angle(x)) @staticmethod def backward(ctx, grad_output): # 无导数操作:反向传播时梯度为0,即该操作不参与上游参数的更新 return torch.zeros_like(grad_output) # 如果需要近似梯度(比如直通估计),也可以返回grad_output,直接把梯度传递过去 # 在网络中使用 class NetEn(nn.Module): def __init__(self): super(NetEn, self).__init__() self.conv1 = nn.Conv2d(1, 64, 5, 2, 0) def forward(self, x): x = self.conv1(x) x = AbsCut.apply(x) return x
这样在反向传播时,这个振幅裁剪操作就不会对上游的conv1参数产生梯度影响,符合你“无导数操作”的需求。
内容的提问来源于stack exchange,提问作者modu01

