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

执行Layer Wise Relevance Propagation时矩阵相乘维度不匹配报错咨询

报错原因解释

该报错的核心原因是LRP实现的层遍历逻辑遗漏了模型中的展平操作,导致全连接层收到的输入维度不匹配:

  • 你的CNN模型forward方法中,经过最后一层卷积池化组layer3后,存在一行out = out.view(out.size(0),-1)的展平逻辑,将4维的特征图(形状为[batch_size, 64, 3, 3])转换为2维的向量(形状为[batch_size, 576]),再喂给全连接层fc1。
  • 但view操作不属于PyTorch的nn.Module子类,你编写的layers = [module for module in model.modules() if not isinstance(module, torch.nn.Sequential)][1:]逻辑无法捕获到该操作,导致特征图在未展平的状态下直接输入全连接层fc1。
  • 报错信息中的mat2是fc1的权重,形状为[10, 576],要求输入的最后一维为576;而未展平的特征图最后一维为3,因此触发矩阵乘法维度不匹配错误。
可行解决方案

提供两种可直接落地的修复方案,任选其一即可:

方案1:修改模型定义,将展平操作封装为标准模块

将原生的view操作替换为PyTorch内置的nn.Flatten模块,使其可以被model.modules()遍历到,无需修改LRP代码:

class Cnn(nn.Module):
    def __init__(self):
       super(Cnn,self).__init__()
    
        self.layer1 = nn.Sequential(
         nn.Conv2d(3,16,kernel_size=3, padding=0,stride=2),
         nn.BatchNorm2d(16),
         nn.ReLU(),
         nn.MaxPool2d(2)
       )
    
        self.layer2 = nn.Sequential(
          nn.Conv2d(16,32, kernel_size=3, padding=0, stride=2),
          nn.BatchNorm2d(32),
          nn.ReLU(),
          nn.MaxPool2d(2)
        )
    
        self.layer3 = nn.Sequential(
          nn.Conv2d(32,64, kernel_size=3, padding=0, stride=2),
          nn.BatchNorm2d(64),
          nn.ReLU(),
          nn.MaxPool2d(2)
       )
        # 新增Flatten模块
        self.flatten = nn.Flatten()
        self.fc1 = nn.Linear(3*3*64,10)
        self.fc2 = nn.Linear(10,2)
        self.relu = nn.ReLU()
    
    def forward(self,x):
       out = self.layer1(x)
       out = self.layer2(out)
       out = self.layer3(out)
       # 替换原来的view操作为调用Flatten模块
       out = self.flatten(out)
       out = self.relu(self.fc1(out))
       out = self.fc2(out)
       return out

修改后原LRP代码可直接正常运行。

方案2:修改LRP代码,手动添加展平逻辑

如果不想修改已训练好的模型结构,可直接在LRP的前向传播循环中,检测到最后一层池化层输出后手动做展平:

def LRP_individual(model, X, device):
   # 获取网络层列表
   layers = [module for module in model.modules() if not isinstance(module, torch.nn.Sequential)][1:]

   # 前向传播输入
   L = len(layers)
   A = [X] + [X] * L # 创建列表存储每一层输出的激活值

   for layer in range(L):
       A[layer + 1] = layers[layer].forward(A[layer])
       # 新增:第三个池化层(索引为11,可打印layers列表确认对应索引)输出后做展平
       if isinstance(layers[layer], nn.MaxPool2d) and layer == 11:
           A[layer+1] = A[layer+1].flatten(start_dim=1)
    
  # LRP函数其余代码

内容的提问来源于stack exchange,提问作者user16738926

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.07 01:15:03