添加第二层GAT层触发矩阵形状不匹配错误,求解决方案
解决多层GAT堆叠时的矩阵维度不匹配问题
使用预训练ResNet18提取特征,尝试添加第二层GAT层时触发错误:mat1 and mat2 shapes cannot be multiplied (64x1 and 512x64),仅用一层GAT时运行正常。
模型代码如下:
class AntispoofModel(nn.Module): def __init__(self, device="cpu", **kwargs): super().__init__() resnet = torch.hub.load('pytorch/vision:v0.10.0', 'resnet18', pretrained=True) self.resnet = nn.Sequential(*[i for i in list(resnet.children())[:-2]]).to(device) for ch in self.resnet.children(): for param in ch.parameters(): param.requires_grad = False self.gat = GAT(**kwargs).to(device) self.device = device self.adj = torch.tensor(grid_to_graph(3, 3, return_as=np.ndarray)).to(device) def forward(self, x): x = self.resnet(x.to(self.device)) x = nn.functional.avg_pool2d(x, 2) x = x.view(-1, 9, 512) x = self.gat(x, self.adj) x = self.gat(x, self.adj) return torch.sigmoid(x)
已尝试调整数据预处理和GAT初始化参数,但仅改变矩阵数值,问题未解决:
transforms.Resize(256), transforms.CenterCrop(224), transforms.ToTensor(), transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]), ]) torch.manual_seed(42) kwargs = {"nfeat":512, "nhid":64, "nclass":1, "nheads":9, "dropout":0.6, "alpha":0.01} train_dataloader = DataLoader(train_dataset, batch_size=64, shuffle=True) test_dataloader = DataLoader(test_dataset, batch_size=64, shuffle=True)
问题根源
当前复用的GAT模块,第一次调用后输出维度已变为nclass=1(从kwargs配置可知),但第二次调用时该模块仍期望输入特征维度为nfeat=512,导致矩阵乘法维度不匹配。
修正方案
1. 拆分GAT层,适配中间维度
定义两个独立的GAT层:第一层将512维输入特征映射到64维(nhid),第二层将64维特征映射到1维输出(nclass)。修改后的模型代码:
class AntispoofModel(nn.Module): def __init__(self, device="cpu", **kwargs): super().__init__() resnet = torch.hub.load('pytorch/vision:v0.10.0', 'resnet18', pretrained=True) self.resnet = nn.Sequential(*[i for i in list(resnet.children())[:-2]]).to(device) for ch in self.resnet.children(): for param in ch.parameters(): param.requires_grad = False # 第一层GAT:输入512维,输出64维 self.gat1 = GAT(nfeat=kwargs["nfeat"], nhid=kwargs["nhid"], nclass=kwargs["nhid"], nheads=kwargs["nheads"], dropout=kwargs["dropout"], alpha=kwargs["alpha"]).to(device) # 第二层GAT:输入64维,输出1维 self.gat2 = GAT(nfeat=kwargs["nhid"], nhid=kwargs["nhid"], nclass=kwargs["nclass"], nheads=kwargs["nheads"], dropout=kwargs["dropout"], alpha=kwargs["alpha"]).to(device) self.device = device self.adj = torch.tensor(grid_to_graph(3, 3, return_as=np.ndarray)).to(device) def forward(self, x): x = self.resnet(x.to(self.device)) x = nn.functional.avg_pool2d(x, 2) x = x.view(-1, 9, 512) x = self.gat1(x, self.adj) x = self.gat2(x, self.adj) return torch.sigmoid(x)
2. 自定义GAT模块的输出逻辑(若使用自定义实现)
如果是自己实现的GAT类,可在模块中增加参数控制输出维度,比如添加is_output_layer标识,让中间层输出nhid维度,输出层输出nclass维度,避免重复定义多个GAT实例。
关键说明
- 原代码复用同一个
self.gat,第一次调用后输出维度变为1,第二次调用时输入维度1与GAT期望的512不匹配,触发报错。 - 拆分GAT层后,每层输入输出维度严格对应:512→64→1,满足矩阵乘法的维度要求。
内容的提问来源于stack exchange,提问作者123321
相关产品推荐
相关产品推荐

