如何通过统一调用获取Tensor与NumPy数组的形状?
统一获取Tensor与NumPy数组形状的可行方案
当然可行,有几种简单且通用的方法可以实现:
方法1:直接访问.shape属性
NumPy数组、PyTorch Tensor、TensorFlow Tensor等主流张量类型都内置了.shape属性,这是最直接的统一调用方式:
import numpy as np import torch import tensorflow as tf def get_shape(obj): # 若需要统一返回元组格式(PyTorch的shape是torch.Size对象,属于元组子类) return tuple(obj.shape) # 测试用例 np_arr = np.array([[0, 1], [2, 3]]) print(get_shape(np_arr)) # 输出:(2, 2) torch_tensor = torch.tensor([[0, 1], [2, 3]]) print(get_shape(torch_tensor)) # 输出:(2, 2) tf_tensor = tf.constant([[0, 1], [2, 3]]) print(get_shape(tf_tensor)) # 输出:(2, 2)
方法2:使用np.shape()函数
NumPy的shape()函数支持直接传入各类Tensor对象,返回标准的元组形状,无需额外转换:
import numpy as np import torch def get_shape(obj): return np.shape(obj) # 测试用例 torch_tensor = torch.randn(3, 5) print(get_shape(torch_tensor)) # 输出:(3, 5) np_arr = np.random.rand(3, 5) print(get_shape(np_arr)) # 输出:(3, 5)
注意事项
- 上述两种方法对JAX Tensor等其他主流张量类型同样适用,兼容性很强;
- 如果需要处理一些小众张量类型,可以额外添加类型判断兜底,但绝大多数场景下上述两种方法足够覆盖需求。
内容的提问来源于stack exchange,提问作者brando f
相关产品推荐
相关产品推荐

