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

基于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

二、通过损失约束实现(适合无法共享参数的场景)

如果因特殊需求不能直接共享参数,可以通过添加约束损失强制卷积核对齐,但要避免之前的均值坍缩问题:

  1. 选择一个固定基准核(比如第一个卷积层的核),让其他核向它靠拢,而非向均值对齐
  2. 给约束损失设置合理权重,平衡任务损失与约束强度

损失计算示例

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.23 00:50:53