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

如何区分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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.29 16:00:04