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

PyTorch Geometric异构图回归模型报错:mat1与mat2 dtype需一致

问题描述

构建基于PyTorch Geometric的异构图回归模型时,运行代码出现错误:RuntimeError: mat1 and mat2 must have the same dtype。forward方法中打印x的数据类型显示为Proxy(getattr_1),相关代码如下:

import torch.nn.functional as F
import torch_geometric.transforms as T
from torch_geometric.nn import SAGEConv, to_hetero
from torch_geometric.nn import global_mean_pool
from torch_geometric.nn import Linear, SAGEConv, to_hetero


class GNNHetero(torch.nn.Module):
    def __init__(self, hidden_channels):
        super().__init__()
        self.conv1 = SAGEConv((-1, -1), hidden_channels)
        self.conv2 = SAGEConv((-1, -1), 1)

    def forward(self, x, edge_index):
        x = self.conv1(x, edge_index).relu()
        x = self.conv2(x, edge_index)
        print(x.dtype)

        return x
    
data = dataset[0]
model = to_hetero(GNNHetero(64), data.metadata(), aggr='sum')

from torch.utils.tensorboard import SummaryWriter
writer = SummaryWriter()
optimizer = torch.optim.Adam(model.parameters(), lr=0.001)
criterion = torch.nn.MSELoss()

def train_hetero():
    model.train()

    for batch in train_loader:  # Iterate in batches over the training dataset.

        out = model(batch.x_dict, batch.edge_index_dict)  # Perform a single forward pass.
        target = data.y.unsqueeze(1)
        loss = criterion(out, target)  # Compute the loss.
        loss.backward()  # Derive gradients.
        optimizer.step()  # Update parameters based on gradients.
        optimizer.zero_grad()  # Clear gradients.

for epoch in range(1, 171):
    print(f'Epoch: {epoch}')
    train_hetero()

print('Done!')

注:dataset是包含1000个HeteroData对象的列表。

解决建议
  • 统一数据类型:错误根源是矩阵运算时输入张量与模型参数 dtype 不匹配。PyTorch Geometric 模型参数默认是torch.float32,需确保数据集里所有节点特征和标签的 dtype 与之一致:
    # 遍历数据集统一所有节点特征的类型
    for data in dataset:
        for node_type in data.x_dict:
            data.x_dict[node_type] = data.x_dict[node_type].float()
        # 统一标签y的类型
        data.y = data.y.float()
    
  • 修正训练循环中的标签引用错误:当前代码里target = data.y.unsqueeze(1)使用的是dataset[0]的标签,而非当前batch的标签,应修改为:
    target = batch.y.unsqueeze(1).float()
    
  • 处理异构模型的输出字典:to_hetero包装后的模型输出是一个字典,对应不同节点类型的预测结果。如果回归任务针对特定节点类型,需从字典中取出对应部分,比如目标节点类型为'target_node':
    out = model(batch.x_dict, batch.edge_index_dict)['target_node']
    
  • 显式转换forward中的张量类型:Proxy(getattr_1)是PyTorch Geometric异构模块内部的代理类型,可在forward中显式将输入转为指定类型,避免 dtype 不匹配:
    def forward(self, x, edge_index):
        x = x.float()  # 强制转为float32,与模型参数一致
        x = self.conv1(x, edge_index).relu()
        x = self.conv2(x, edge_index)
        return x
    

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.17 10:37:37