PyTorch调用mean()触发Got Long报错,int64被识别为Long的原因是什么
问题解答
在PyTorch的类型体系中,Long就是int64类型的别名,二者完全等价,不存在识别错误的问题。
具体逻辑说明
- 你传入全为整数的Python列表
[1,0,1]创建张量时,PyTorch默认会将整数序列推断为int64(即报错信息中提到的Long)类型,这是框架固定的默认类型推断规则。 mean()这类需要输出浮点结果的算子,设计时就不支持直接对整型张量计算,目的是避免出现用户预期外的精度丢失,因此会明确要求输入为浮点或复数类型,这是报错的根本触发原因。- 你修改为包含浮点值的列表
[1.,0,1]后,PyTorch会自动将张量推断为浮点类型,符合mean()算子的输入要求,因此可以正常返回结果。
额外解决方案
如果不想修改原始列表的数值写法,也可以在创建张量时手动指定数据类型:
tns = torch.tensor([1,0,1], dtype=torch.float32) tns.mean() # 可正常输出 tensor(0.6667)
内容的提问来源于stack exchange,提问作者Chris
相关产品推荐
相关产品推荐

