如何通过断言检查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')), "张量中存在正无穷或非法非有限值"
原理说明
- 提取非有限元素:通过
~mask取反掩码,获取所有非有限值的位置,再用t[~mask]提取这些元素 - 批量检查:
torch.all(t[~mask] == float('-inf'))会检查所有提取出的元素是否全等于负无穷,返回布尔值 - 边界情况处理:如果张量中没有非有限值,
t[~mask]会是空张量,此时torch.all默认返回True,断言依然通过(因为没有违反条件的元素)
内容的提问来源于stack exchange,提问作者melatonin15
相关产品推荐
相关产品推荐

