Python中PyTorch张量乘法返回inf,如何改用float128/Decimal?
解决PyTorch回归运算返回inf张量的类型处理方法
使用PyTorch的float128类型
直接将参与运算的所有张量转换为torch.float128类型,利用其更大的数值范围避免溢出。注意torch.float128主要支持CPU环境,多数GPU暂不兼容:
import torch def regression(my_x, my_m, my_b): # 转换所有输入张量为float128 x_128 = my_x.to(torch.float128) m_128 = my_m.to(torch.float128) b_128 = my_b.to(torch.float128) return m_128 * x_128 + b_128
如果是GPU环境,可尝试torch.bfloat16(精度略低于float128,但GPU支持更好,仅当溢出由精度不足而非数值范围导致时适用)。
使用Python Decimal类型
Decimal提供任意精度的十进制运算,但需要将PyTorch张量转换为Python数值类型,会脱离PyTorch计算图,适合小规模非梯度依赖场景:
import torch from decimal import Decimal, getcontext # 设置Decimal的计算精度(默认28位,可按需调高) getcontext().prec = 50 def regression(my_x, my_m, my_b): # 将张量元素转换为Decimal对象 x_dec = [Decimal(str(val.item())) for val in my_x.flatten()] m_dec = [Decimal(str(val.item())) for val in my_m.flatten()] b_dec = [Decimal(str(val.item())) for val in my_b.flatten()] # 执行回归运算 result_dec = [m * x + b for x, m, b in zip(x_dec, m_dec, b_dec)] # 可选:转回PyTorch张量(需转成float64避免精度丢失) result_tensor = torch.tensor([float(val) for val in result_dec], dtype=torch.float64).reshape(my_x.shape) return result_tensor
额外注意
- 先验证溢出原因:打印输入张量的最大值(
print(my_x.max(), my_m.max())),确认是否接近float32上限(约3.4e38),再针对性处理。 - float128运算速度慢于float32/64,需权衡精度与性能。
内容的提问来源于stack exchange,提问作者Pritesh
相关产品推荐
相关产品推荐

