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

添加第二层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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.30 20:45:40