基于PyTorch构建损失函数使所有卷积核完全相同
让CNN所有卷积核保持一致的两种实现方法
一、直接共享卷积核参数(最优方案)
这是最直接且高效的方式,不需要额外损失约束——让所有目标卷积层共用同一个可学习参数张量,反向传播时所有层的梯度都会累积到这个共享参数上,自然保证所有卷积核完全一致,还能避免参数坍缩问题。
针对你描述的3层单通道CNN示例代码
import torch import torch.nn as nn class SharedKernelNet(nn.Module): def __init__(self): super().__init__() # 定义共享的2×2×1卷积核参数 self.shared_kernel = nn.Parameter(torch.randn(1, 1, 2, 2)) # 初始化卷积层,替换weight为共享参数 self.conv1 = nn.Conv2d(1, 1, kernel_size=2, bias=True) self.conv2 = nn.Conv2d(1, 1, kernel_size=2, bias=True) self.conv3 = nn.Conv2d(1, 1, kernel_size=2, bias=True) self.conv1.weight = self.shared_kernel self.conv2.weight = self.shared_kernel self.conv3.weight = self.shared_kernel # 后续层按需添加 self.relu = nn.ReLU() self.pool = nn.MaxPool2d(2) def forward(self, x): x = self.relu(self.conv1(x)) x = self.pool(x) x = self.relu(self.conv2(x)) x = self.pool(x) x = self.relu(self.conv3(x)) # 全连接层等后续操作 return x
适配你提供的MNIST多通道CNN代码
如果要让每层内的所有输出通道卷积核保持一致(比如第一层16个输出通道共用同一个3×3×1核),可以修改模型如下:
import torch import torchvision from tqdm import tqdm import matplotlib class Net(torch.nn.Module): def __init__(self): super(Net,self).__init__() # 定义各层的共享卷积核参数 self.shared_k1 = nn.Parameter(torch.randn(1, 1, 3, 3)) self.shared_k2 = nn.Parameter(torch.randn(1, 16, 3, 3)) self.shared_k3 = nn.Parameter(torch.randn(1, 32, 3, 3)) self.model = torch.nn.Sequential( torch.nn.Conv2d(in_channels = 1,out_channels = 16,kernel_size = 3,stride = 1,padding = 1), torch.nn.ReLU(), torch.nn.MaxPool2d(kernel_size = 2,stride = 2), torch.nn.Conv2d(in_channels = 16,out_channels = 32,kernel_size = 3,stride = 1,padding = 1), torch.nn.ReLU(), torch.nn.MaxPool2d(kernel_size = 2,stride = 2), torch.nn.Conv2d(in_channels = 32,out_channels = 64,kernel_size = 3,stride = 1,padding = 1), torch.nn.ReLU(), torch.nn.Flatten(), torch.nn.Linear(in_features = 7 * 7 * 64,out_features = 128), torch.nn.ReLU(), torch.nn.Linear(in_features = 128,out_features = 10), torch.nn.Softmax(dim=1) ) # 将共享核重复对应次数,替换卷积层weight self.model[0].weight = nn.Parameter(self.shared_k1.repeat(16, 1, 1, 1)) self.model[3].weight = nn.Parameter(self.shared_k2.repeat(32, 1, 1, 1)) self.model[6].weight = nn.Parameter(self.shared_k3.repeat(64, 1, 1, 1)) def forward(self,input): output = self.model(input) return output
二、通过损失约束实现(适合无法共享参数的场景)
如果因特殊需求不能直接共享参数,可以通过添加约束损失强制卷积核对齐,但要避免之前的均值坍缩问题:
- 选择一个固定基准核(比如第一个卷积层的核),让其他核向它靠拢,而非向均值对齐
- 给约束损失设置合理权重,平衡任务损失与约束强度
损失计算示例
criterion_task = torch.nn.CrossEntropyLoss() lambda_reg = 0.1 # 约束损失权重,需根据任务调整 # 计算任务损失 outputs = net(trainImgs) loss_task = criterion_task(outputs, labels) # 计算约束损失:让conv2、conv3的核与conv1对齐 loss_reg = torch.norm(net.model[3].weight - net.model[0].weight, p=2) + torch.norm(net.model[6].weight - net.model[0].weight, p=2) # 总损失 total_loss = loss_task + lambda_reg * loss_reg # 反向传播与优化 total_loss.backward() optimizer.step()
为什么用均值约束会导致核趋近于0?
当以所有核的均值为目标时,优化目标是sum((k_i - mean_k)^2) + 任务损失。这个约束项的最小值可在所有k_i相等时达成,而如果任务损失在k_i全为0时也能取得局部最优(比如MNIST任务中全0核输出可由全连接层弥补),优化器就会倾向于让所有核坍缩到0。而固定基准核的方式,约束项的最小值是其他核等于基准核,基准核会被任务损失驱动更新,不会陷入无意义的0值局部最优。
内容的提问来源于stack exchange,提问作者Liang Sun
相关产品推荐
相关产品推荐

