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

