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

PyTorch反向传播报错:检测到Double类型但预期为Float类型

解决方案

问题根源

  1. PyTorch Geometric(PyG)中edge_index必须是整数类型(如torch.long),你的edge_index是torch.float64,本身不符合GNN层的输入要求,会引发后续dtype混乱。
  2. 模型参数已转为torch.float32,但目标值graph.Utility仍为torch.float64,导致loss的dtype与模型参数dtype不匹配,反向传播时触发类型错误。

具体修复步骤

方案一:统一使用float32(推荐,多数场景下显存占用更低、计算更快)

# 修正edge_index的类型(必须步骤)
graph.edge_index = graph.edge_index.long()

# 将所有输入数据和目标值转为float32,与模型保持一致
graph.x = graph.x.to(torch.float32)
graph.Utility = graph.Utility.to(torch.float32)

# 确保模型参数为float32(你已执行过,可再次确认)
model = model.to(torch.float32)

# 重新计算loss并反向传播
loss = F.mse_loss(model(graph.x, graph.edge_index, graph.batch), graph.Utility)
loss.backward()

方案二:统一使用float64(适用于需要高精度的场景)

# 修正edge_index的类型(必须步骤)
graph.edge_index = graph.edge_index.long()

# 将模型转为float64,与现有数据dtype保持一致
model = model.to(torch.float64)

# 重新计算loss并反向传播
loss = F.mse_loss(model(graph.x, graph.edge_index, graph.batch), graph.Utility)
loss.backward()

关键说明

  • 无论哪种方案,修正edge_index的整数类型是核心前提,浮点型的节点索引会导致GNN层内部计算逻辑出错,进而引发dtype不匹配问题。
  • 模型、输入数据、目标值三者的dtype必须完全统一,否则反向传播时梯度计算会触发类型校验错误。

内容的提问来源于stack exchange,提问作者Tony Sirico

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.09 17:55:10