pytorch torch.nan_to_num函数介绍

torch.nan_to_num 函数简介

torch.nan_to_num 是 PyTorch 中的一个函数,用于将张量中的特殊浮点值(如 NaN、+Inf 和 -Inf)替换为指定的数值,或使用默认替代值。

函数签名

复制代码
torch.nan_to_num(input, nan=0.0, posinf=None, neginf=None)

参数

  1. input:

    • 输入张量。
    • 可以包含 NaN、正无穷(+Inf)、负无穷(-Inf)等特殊值。
  2. nan (可选):

    • 替换 NaN 的值。
    • 默认是 0.0。
  3. posinf (可选):

    • 替换正无穷 (+Inf) 的值。
    • 默认是张量元素的最大有限值 (torch.finfo(input.dtype).max)。
  4. neginf (可选):

    • 替换负无穷 (-Inf) 的值。
    • 默认是张量元素的最小有限值 (torch.finfo(input.dtype).min)。

返回值

  • 返回一个张量,其中的 NaN、+Inf 和 -Inf 被替换为指定的值。
  • 输出张量与输入张量的形状和数据类型相同。

工作原理

  • NaN : 检测到 NaN 后,替换为参数 nan 指定的值。
  • +Inf 和 -Inf : 检测到无穷值后,分别替换为参数 posinf 和 neginf 指定的值。

简单示例

复制代码
import torch

# 创建包含 NaN、+Inf 和 -Inf 的张量
x = torch.tensor([float('nan'), float('inf'), -float('inf'), 1.0, -2.0])

# 替换 NaN 和 Inf
result = torch.nan_to_num(x, nan=0.0, posinf=10.0, neginf=-10.0)
print(result)

输出:

复制代码
tensor([  0.,  10., -10.,   1.,  -2.])

使用默认值

如果没有指定 posinf 和 neginf,函数会使用数据类型的最大或最小值。

复制代码
x = torch.tensor([float('nan'), float('inf'), -float('inf')], dtype=torch.float32)

result = torch.nan_to_num(x)
print(result)

输出:

复制代码
tensor([ 0.0000e+00,  3.4028e+38, -3.4028e+38])

其中 3.4028e+38 和 -3.4028e+38 分别是 float32 类型的最大和最小有限值。

广播支持

torch.nan_to_num 支持广播机制,当输入包含多维张量时同样可以逐元素替换:

复制代码
x = torch.tensor([[float('nan'), float('inf')], [-float('inf'), 1.0]])
result = torch.nan_to_num(x, nan=0.0, posinf=100.0, neginf=-100.0)
print(result)

输出:

复制代码
tensor([[   0.,  100.],
        [-100.,    1.]])

应用场景

1. 清洗数据 : 替换缺失值(NaN)或异常值(+Inf、-Inf)以便进一步处理。

复制代码
x = torch.tensor([float('nan'), 5.0, float('inf'), -float('inf')])
clean_x = torch.nan_to_num(x, nan=0.0)
print(clean_x)  # tensor([ 0.,  5.,  max_value, min_value])

2. 防止计算异常 : 在模型训练或推理过程中,防止出现 NaN 或无穷值导致的计算失败。

3. 图像/信号处理: 在处理图像或信号数据时,用于替换缺失的像素值或异常值。

注意事项

  1. 数据类型兼容性:

    • 如果输入张量的类型为整数,使用 torch.nan_to_num 会报错,因为整数类型无法表示 NaN 或无穷值。
    • 函数只能用于浮点类型张量(如 torch.float32, torch.float64)。
  2. 默认替换值:

    • 对于正无穷和负无穷,默认替换值依赖于张量的数据类型。
  3. 性能开销:

    • 对大张量来说,函数调用会带来一定的计算开销,需在实际应用中注意性能。

总结

torch.nan_to_num 是处理数据异常(如缺失值和溢出值)的重要工具,特别适用于数据预处理和深度学习模型的训练过程。通过灵活的参数设置,可以有效替换各种特殊值,保证后续计算的稳定性和可靠性。

相关推荐
lank_M8 小时前
截图OCR预处理在普通屏上翻车,Retina截图却没事
图像处理·人工智能·计算机视觉·ocr
Aloudata8 小时前
LookML 语义模型 vs 企业级独立语义层:BI 建模语言能否承担企业语义底座?
数据库·人工智能·数据分析·数据资产·dataagent
启观川8 小时前
数据结构与算法 -第 3 章 常用算法-动态规划
数据结构·笔记·python·算法
码爸8 小时前
排序算法介绍
python·算法·排序算法
龙亘川8 小时前
AI 协同赋能城市治理:支撑政协数字化履职的技术路径探析
大数据·人工智能·智慧城市·开源软件·数据可视化
通信大模型8 小时前
IEEE TCCN | 面向低空经济网络的Agentic AI驱动多无人机轨迹优化
网络·人工智能·无人机
笔墨登场说说8 小时前
flink bin/start-cluster.sh 帮我做成开机启动
开发语言·python
爱编程的小白L8 小时前
2027 计算机毕业设计选题汇总|深度学习专项(2027最新)
人工智能·深度学习·课程设计
逸模8 小时前
BIM在连锁餐饮装修中的应用:不只是画三维图
大数据·数据库·人工智能·物联网·建模