使用分组卷积实现重复卷积时的kernel反向传播问题咨询
最优实现方案
直接通过动态生成分组卷积权重的方式,让计算图直接关联原始单通道kernel,反向传播时梯度自动聚合到kernel参数,全程无额外后处理开销,是速度最优的实现:
import torch import torch.nn as nn torch.manual_seed(0) class SharedKernelGroupConv(nn.Module): def __init__(self, in_channels, kernel_size, bias=True, **conv_kwargs): super().__init__() # 仅维护单个原始卷积核,这是唯一需要更新的参数 self.single_kernel = nn.Conv2d(1, 1, kernel_size, bias=bias, **conv_kwargs) self.in_channels = in_channels self.conv_kwargs = conv_kwargs def forward(self, x): # 前向时动态生成分组卷积需要的权重、偏置,计算图直接绑定原始kernel参数 weight = self.single_kernel.weight.repeat(self.in_channels, 1, 1, 1) bias = self.single_kernel.bias.repeat(self.in_channels) if self.single_kernel.bias is not None else None return nn.functional.conv2d( x, weight, bias, groups=self.in_channels, **self.conv_kwargs ) # 验证用例 x = torch.rand(1, 3, 3, 3) first = x[:, 0:1, ...] second = x[:, 1:2, ...] third = x[:, 2:3, ...] conv = SharedKernelGroupConv(in_channels=3, kernel_size=3, padding=0) # 前向结果验证,和你原示例输出完全一致 print(conv(x)) print(conv.single_kernel(first), conv.single_kernel(second), conv.single_kernel(third)) # 反向传播验证 loss = conv(x).sum() loss.backward() # 梯度会直接聚合到single_kernel的参数上,不需要额外处理 print(conv.single_kernel.weight.grad) print(conv.single_kernel.bias.grad) # 优化器仅需传入single_kernel的参数即可 optimizer = torch.optim.SGD(conv.single_kernel.parameters(), lr=1e-3) optimizer.step()
方案说明
- 所有计算都复用pytorch原生优化的卷积算子,仅增加了一次权重
repeat操作,开销可以忽略不计 - 反向传播时,分组卷积三个通道的权重梯度会自动累加回原始的
single_kernel参数,完全符合你仅需要更新单个kernel的需求 - 不需要手动维护分组卷积层和原始kernel的参数同步,避免了训练过程中参数不一致的问题
内容的提问来源于stack exchange,提问作者S P Sharan
相关产品推荐
相关产品推荐

