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

PyTorch自定义激活函数:类实现与函数实现的差异探究

PyTorch自定义激活函数:普通函数vs.autograd.Function类的差异与无导数操作实现

我在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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.08 02:47:06