如何在PyTorch中为分组卷积每组跨所有通道复用单一权重?
实现分组内权重共享的2D卷积层
你需要的是在分组卷积中让每个分组内的所有输出通道共享同一组权重,以此减少参数数量缓解过拟合,以下是PyTorch中的具体实现方案:
自定义共享权重的分组卷积层
import torch import torch.nn as nn import torch.nn.functional as F class Conv2dSharedGroups(nn.Module): def __init__(self, in_channels, out_channels, kernel_size, stride=1, groups=1, bias=False): super().__init__() # 验证通道数与分组数的整除性 assert in_channels % groups == 0, "输入通道数必须能被分组数整除" assert out_channels % groups == 0, "输出通道数必须能被分组数整除" self.in_channels = in_channels self.out_channels = out_channels self.kernel_size = kernel_size self.stride = stride self.groups = groups self.per_group_in = in_channels // groups # 每组输入通道数:30//5=6 self.per_group_out = out_channels // groups # 每组输出通道数:30//5=6 # 定义共享权重:每组对应1个卷积核,形状为(groups, 1, kernel_size[0], kernel_size[1]) self.weight = nn.Parameter(torch.randn(groups, 1, *kernel_size)) if bias: self.bias = nn.Parameter(torch.zeros(out_channels)) else: self.bias = None def forward(self, x): # 扩展共享权重到标准分组卷积的权重形状:(out_channels, per_group_in, kernel_size[0], kernel_size[1]) # 1. 先在输出通道维度重复每组的权重,得到(30, 1, 1, 20) expanded_weight = self.weight.repeat_interleave(self.per_group_out, dim=0) # 2. 再在输入通道维度重复,得到(30, 6, 1, 20) expanded_weight = expanded_weight.repeat_interleave(self.per_group_in, dim=1) # 调用标准卷积函数,使用分组参数 return F.conv2d( x, expanded_weight, bias=self.bias, stride=self.stride, groups=self.groups )
代码说明
- 参数精简:仅定义
groups个卷积核(形状(5,1,1,20)),相比原卷积层的30×6×1×20权重,参数数量从3600减少到100,大幅降低模型复杂度。 - 权重扩展逻辑:在
forward中,将每组的单一权重分别在输出通道维度重复6次(对应每组6个输出通道)、输入通道维度重复6次(对应每组6个输入通道),扩展为标准分组卷积所需的权重形状,再调用F.conv2d完成计算。 - 接口兼容性:保持与原
nn.Conv2d相同的输入输出接口,可直接替换原卷积层使用。
使用示例
替换你原来的卷积层:
# 原卷积层 # original_conv = nn.Conv2d(kernel_size=(1,20), stride=1, groups=5, out_channels=30, in_channels=30, bias=False) # 替换为自定义共享权重卷积层 shared_conv = Conv2dSharedGroups( in_channels=30, out_channels=30, kernel_size=(1,20), stride=1, groups=5, bias=False )
内容的提问来源于stack exchange,提问作者Anonymous
相关产品推荐
相关产品推荐

