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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.10 07:45:26