执行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
相关产品推荐
相关产品推荐

