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

如何在PyTorch中实现Conv2d的反向传播函数?

实现Conv2d反向传播的步骤与代码

先理清Linear函数的实现逻辑

PyTorch自定义自动求导逻辑的核心是继承torch.autograd.Function,并实现静态的forward和backward方法,再通过nn.Module封装成可复用的层:

  • LinearFunction
    • forward:接收输入、权重、偏置,用ctx.save_for_backward保存反向传播需要的张量,执行矩阵乘法+偏置广播相加得到输出。
    • backward:接收上游梯度grad_output,根据ctx.needs_input_grad判断哪些输入需要计算梯度:
      • 输入梯度:grad_output与权重做矩阵乘法
      • 权重梯度:grad_output转置后与输入做矩阵乘法
      • 偏置梯度:对grad_output在样本维度求和
  • Linear模块:继承nn.Module,初始化可学习参数(nn.Parameter),forward方法调用LinearFunction.apply完成前向计算。

Conv2d反向传播实现

Conv2d的反向传播需要计算输入、卷积核、偏置三者的梯度,以下是完整实现:

1. 自定义Conv2dFunction

import torch
import torch.nn as nn
import torch.nn.functional as F
from torch.autograd import Function

class Conv2dFunction(Function):
    @staticmethod
    def forward(ctx, input, weight, bias=None, stride=1, padding=0, dilation=1, groups=1):
        # 保存反向传播所需的张量和卷积参数
        ctx.save_for_backward(input, weight, bias)
        ctx.stride = stride
        ctx.padding = padding
        ctx.dilation = dilation
        ctx.groups = groups
        
        # 复用PyTorch官方Conv2d前向计算逻辑
        output = F.conv2d(input, weight, bias, stride, padding, dilation, groups)
        return output

    @staticmethod
    def backward(ctx, grad_output):
        input, weight, bias = ctx.saved_tensors
        grad_input = grad_weight = grad_bias = None
        
        # 计算输入的梯度:通过转置卷积实现
        if ctx.needs_input_grad[0]:
            grad_input = F.conv_transpose2d(
                grad_output, weight, None, 
                stride=ctx.stride, padding=ctx.padding, 
                dilation=ctx.dilation, groups=ctx.groups
            )
        
        # 计算卷积核的梯度:对填充后的输入与上游梯度做特定参数的卷积
        if ctx.needs_input_grad[1]:
            input_padded = F.pad(input, (ctx.padding, ctx.padding, ctx.padding, ctx.padding))
            grad_weight = F.conv2d(
                input_padded.transpose(0, 1), 
                grad_output.transpose(0, 1), 
                None, 
                stride=ctx.dilation, padding=0, 
                dilation=ctx.stride, groups=ctx.groups
            ).transpose(0, 1)
            # 确保梯度形状与原权重一致
            grad_weight = grad_weight[:weight.shape[0], :weight.shape[1], :, :]
        
        # 计算偏置的梯度:对上游梯度在样本、高度、宽度维度求和
        if bias is not None and ctx.needs_input_grad[2]:
            grad_bias = grad_output.sum((0, 2, 3))
        
        # 返回与forward输入对应的梯度,多余参数返回None
        return grad_input, grad_weight, grad_bias, None, None, None, None

2. 自定义Conv2d模块

class Conv2d(nn.Module):
    def __init__(self, in_channels, out_channels, kernel_size, stride=1, padding=0, dilation=1, groups=1, bias=True):
        super().__init__()
        self.in_channels = in_channels
        self.out_channels = out_channels
        # 统一处理kernel_size、stride等参数为元组格式
        self.kernel_size = kernel_size if isinstance(kernel_size, tuple) else (kernel_size, kernel_size)
        self.stride = stride if isinstance(stride, tuple) else (stride, stride)
        self.padding = padding if isinstance(padding, tuple) else (padding, padding)
        self.dilation = dilation if isinstance(dilation, tuple) else (dilation, dilation)
        self.groups = groups
        
        # 初始化卷积核参数
        self.weight = nn.Parameter(torch.empty(out_channels, in_channels//groups, *self.kernel_size))
        if bias:
            self.bias = nn.Parameter(torch.empty(out_channels))
        else:
            self.register_parameter('bias', None)
        
        # 初始化参数值
        nn.init.kaiming_uniform_(self.weight, mode='fan_in', nonlinearity='relu')
        if self.bias is not None:
            nn.init.zeros_(self.bias)
    
    def forward(self, input):
        return Conv2dFunction.apply(
            input, self.weight, self.bias, 
            self.stride, self.padding, self.dilation, self.groups
        )

关键逻辑说明

  • 前向传播直接复用PyTorch官方F.conv2d保证计算效率,同时保存必要参数供反向使用。
  • 输入梯度通过转置卷积F.conv_transpose2d计算,这是卷积反向传播中输入梯度的标准实现方式。
  • 权重梯度需要对填充后的输入和上游梯度做适配参数的卷积,同时兼容分组卷积的逻辑。
  • 偏置梯度的计算逻辑与Linear层一致,仅需对上游梯度在非通道维度求和。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.06 20:20:36