PyTorch转换NumPy数组为float64张量时如何保留8位以上有效数字?
解决PyTorch张量打印时保留8位以上有效数字的问题
首先明确:你用torch.float64转换后的张量本身已经保留了完整的双精度数据(对应NumPy的float64,可存储约15-17位有效数字),只是PyTorch默认打印设置只显示5位有效数字,以下是几种解决方法:
全局调整PyTorch打印精度
通过torch.set_printoptions()设置全局打印的有效数字位数,之后所有张量打印都会遵循这个设置:import torch import numpy as np numpy_array = np.array([-8.32457799, -8.18170165, -8.03901151, -4.34838355, -4.33105147, -4.31420002]) # 设置打印精度为8位有效数字 torch.set_printoptions(precision=8) tensor = torch.tensor(numpy_array, dtype=torch.float64) print(tensor) print(tensor[0])输出结果:
tensor([-8.32457799, -8.18170165, -8.03901151, -4.34838355, -4.33105147, -4.31420002], dtype=torch.float64) tensor(-8.32457799, dtype=torch.float64)针对单个元素打印完整精度
用.item()方法取出张量元素的Python原生数值,直接打印即可显示完整精度:print(tensor[0].item())输出:
-8.32457799转换为NumPy数组打印
NumPy默认会显示更多有效数字,将张量转回NumPy数组后打印,能直接看到原始精度:print(tensor.numpy())输出与原始NumPy数组精度一致。
内容的提问来源于stack exchange,提问作者asdfe
相关产品推荐
相关产品推荐

