为何PyTorch要求将张量的整数形状值断言为元组?
PyTorch张量shape的统一元组设计问题
PyTorch框架中,所有张量的shape属性固定返回元组类型,这是为了保持接口一致性——不管张量是0维(标量)、1维还是多维,shape的类型统一,避免因维度不同导致的类型判断混乱。
你的代码问题分析
x = torch.Tensor(3,3)创建的是二维张量,shape为(3,3)(元组),所以assert x.shape == (3,3)能通过;y = torch.Tensor(9)创建的是一维张量(长度为9),它的shape是(9,)(单元素元组),而(9)在Python中只是整数9,和元组(9,)类型完全不同,因此assert y.shape == (9)必然失败。
实用解决方法
- 如果你需要获取一维张量的长度,可以直接用:
print(len(y)) # 输出9 print(y.shape[0]) # 输出9 - 如果要创建标量张量(shape为空元组
()),应该用小写的tensor函数传入单个数值:scalar = torch.tensor(9) print(scalar.shape) # 输出() assert scalar.shape == () # 断言成功 - 也可以通过维度数
ndim判断张量类型:assert y.ndim == 1 # 一维张量,断言成功 assert scalar.ndim == 0 # 标量,断言成功
内容的提问来源于stack exchange,提问作者desert_ranger
相关产品推荐
相关产品推荐

