如何区分PyTorch常规Tensor与Nested Tensor?
区分PyTorch常规张量与Nested Tensor的可靠方法
PyTorch推出的Nested Tensor创建后,用type(nt) == torch.Tensor或isinstance(nt, torch.Tensor)判断都会返回True,无法直接通过类型区分它和常规张量。目前可以通过捕获size()方法抛出的RuntimeError来判断,但这种依赖报错的方式担心会因API变更失效,有没有更简便可靠的方法?
创建示例代码
import torch a = torch.randn(20, 128) nt = torch.nested.nested_tensor([a, a], dtype=torch.float32)
当前依赖报错的判断方法
def is_nested_tensor(nt): if not isinstance(nt, torch.Tensor): return False try: # 尝试无参调用size() nt.size() return False except RuntimeError: return True return False
更可靠的判断方案
可以直接使用PyTorch官方提供的 torch.is_nested() 函数,或者访问张量的 is_nested 属性,这两种方式都是原生支持的判断逻辑,完全规避了依赖报错的不稳定问题,也更简洁。
改进后的判断代码
def is_nested_tensor(nt): if not isinstance(nt, torch.Tensor): return False # 方式1:使用torch.is_nested()函数 return torch.is_nested(nt) # 方式2:访问张量的is_nested属性 # return nt.is_nested
说明
这两个API在PyTorch 1.10及以上版本均可用,是官方专门为区分Nested Tensor设计的,相比依赖size()报错的方式,稳定性和可读性都更强,不会因为后续API的调整而失效。
内容的提问来源于stack exchange,提问作者Matt Pitkin
相关产品推荐
相关产品推荐

