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

如何通过断言检查torch.isfinite为False时张量仅含负无穷?

如何断言PyTorch张量中所有非有限值仅为负无穷?

问题描述

我有一个float类型的PyTorch张量t,其中可能包含float('-inf')或float('inf'),张量规模可能较大。我用torch.isfinite生成了布尔掩码mask,标记每个元素是否为有限值。现在需要添加断言检查,确保mask为False的位置,张量t中仅存在float('-inf'),要求不遍历张量,简洁实现。

示例代码:

import torch

t = torch.tensor(
    [[float('-inf'), 1.5, 0.5, 2.0], 
     [5.0, 3.0, float('-inf'), 5.5],
     [0.5, float('-inf'), 3.0, 2.0]]
)

mask = torch.isfinite(t)
print(mask)
>>> tensor([[False,  True,  True,  True],
          [ True,  True, False,  True],
          [ True, False,  True,  True]])

解决方案

利用PyTorch的张量索引和批量判断功能,无需遍历即可完成检查:

import torch

t = torch.tensor(
    [[float('-inf'), 1.5, 0.5, 2.0], 
     [5.0, 3.0, float('-inf'), 5.5],
     [0.5, float('-inf'), 3.0, 2.0]]
)

mask = torch.isfinite(t)
# 断言所有非有限值均为负无穷
assert torch.all(t[~mask] == float('-inf')), "张量中存在正无穷或非法非有限值"

原理说明

  1. 提取非有限元素:通过~mask取反掩码,获取所有非有限值的位置,再用t[~mask]提取这些元素
  2. 批量检查:torch.all(t[~mask] == float('-inf'))会检查所有提取出的元素是否全等于负无穷,返回布尔值
  3. 边界情况处理:如果张量中没有非有限值,t[~mask]会是空张量,此时torch.all默认返回True,断言依然通过(因为没有违反条件的元素)

内容的提问来源于stack exchange,提问作者melatonin15

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.26 07:03:18