PyTorch张量存储内存计算:三种方法结果不同,哪一种正确?
PyTorch张量存储内存计算方法的正确性疑问
import torch import sys a = torch.rand(10) b = torch.rand(100) sys.getsizeof(a) # 72 sys.getsizeof(a.storage()) # 88 a.element_size() * a.nelement() # 40 sys.getsizeof(b) # 72 sys.getsizeof(b.storage()) # 448 b.element_size() * b.nelement() # 400
经搜索发现,人们使用上述三种方法计算PyTorch张量的存储内存,但三种方法返回的结果各不相同。请问哪一种方法是正确的?
三种方法的差异解析
sys.getsizeof(a):返回的是Python对象本身的内存占用,也就是PyTorch张量这个Python包装类的大小,和张量实际存储的数据无关。所以不管张量元素数量多少,结果基本固定(示例中为72,具体数值随Python版本、系统环境略有差异)。sys.getsizeof(a.storage()):返回PyTorch底层存储对象的总内存,包含张量数据本身加上存储对象的额外管理开销。比如示例中a的存储对象返回88,其中40是数据大小,剩余48是存储对象的元数据开销;b的存储对象返回448,400是数据大小,48是固定开销。a.element_size() * a.nelement():直接计算张量实际存储的数据字节数。element_size()返回单个元素的字节数(示例中torch.rand默认生成float32类型,单元素占4字节),nelement()返回总元素数,两者相乘得到的是纯数据的内存占用。
正确方法的选择
如果你的需求是获取张量实际存储数据的内存大小,a.element_size() * a.nelement()是最准确的,因为它仅计算数据本身,排除了Python对象和存储对象的额外开销。
若需要包含底层存储对象的管理开销,可使用sys.getsizeof(a.storage());但仅关注张量数据占用内存时,前者是正确选择。
内容的提问来源于stack exchange,提问作者Simon
相关产品推荐
相关产品推荐

