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

