导读:本期聚焦于韩兆瑞创作的《PyTorch的Storage存储机制是什么?张量底层数据存储原理与常见问题详解》,敬请观看详情。一个张量经过切片、reshape或view操作后,修改新张量的数值,原始张量的数据竟然跟着变了,这背后的原因就藏在PyTorch的Storage存储机制里。Storage是张量底层的一维连续内存块,张量本身只记录形状、步长和偏移量等元信息,多个张量可以指向同一块Storage。本文系统讲解Storage与Tensor的关系、视图操作的内存共享原理、storage_offset与stride的作用方式,并汇总了clone与copy的区别、切片后内存不释放、序列化时存储的保存行为等高频问题,帮助你彻底搞懂PyTorch的底层数据组织方式。

在PyTorch中,一个张量对象其实由两部分组成:一部分是描述性的元信息,比如形状、数据类型、所在设备;另一部分才是真正存放数值的连续内存块,这个内存块就是Storage。理解Storage机制,是搞清楚视图操作、内存共享、切片行为、序列化等一系列问题的基础。很多看起来诡异的现象,比如改一个切片的值原张量跟着变、删掉大张量内存却不释放,本质上都是Storage在起作用。

PyTorch的Storage存储机制是什么?张量底层数据存储原理与常见问题详解

Storage到底是什么:张量背后的数据容器

Storage可以理解为一个一维的、连续的原始数据数组,它只关心两件事:能装多少个元素、每个元素是什么类型。它没有形状的概念,也没有维度的概念。一个形状为(3, 4)的浮点张量,其背后的Storage就是一个能容纳12个float32元素的一维数组,仅此而已。

张量与Storage的关系可以用一句话概括:Tensor是对Storage的一种视角描述。张量持有指向某个Storage的引用,同时记录了shape(每个维度多长)、stride(沿每个维度移动一个元素时,在Storage中要跨过多少个元素)以及storage_offset(从Storage的第几个元素开始读)。张量在创建时并不总是新建Storage,比如通过torch.arange创建会分配新存储,而通过切片得到的新张量则会复用原有Storage。

可以用下面这段代码直观感受Storage的存在:

import torch

t = torch.tensor([[1.0, 2.0, 3.0],
                  [4.0, 5.0, 6.0]])

# 查看张量底层的一维存储(PyTorch 2.x 推荐用 untyped_storage)
storage = t.untyped_storage()
print(storage)          # 输出类似 1.0 2.0 3.0 4.0 5.0 6.0 的一维数据
print(storage.nbytes()) # 底层存储占用的字节数:24(6个float32,每个4字节)

# 老版本API,PyTorch 2.0之后已被弃用,调用会给出警告
# storage = t.storage()

从输出可以看到,无论张量是几维的,Storage里永远是一维的线性数据。维度信息完全由张量自身的元数据来解释,这也是PyTorch能高效实现转置、切片等操作的关键所在。

形状、步长与偏移:张量如何读取Storage

知道Storage是一维数组后,一个自然的问题是:二维甚至更高维的张量,是怎么从一维数组里取数的?答案是三个元数据协同工作。假设要访问元素t[i][j],实际访问的Storage位置就是storage_offset + i * stride[0] + j * stride[1]。stride决定了每个维度上移动一步对应Storage中的跨度,offset决定了起点。

转置操作最能体现这套机制的精妙之处。对张量做t.t()或t.transpose(0, 1),PyTorch不会搬动任何数据,只是交换了stride中两个维度的值,同时shape也跟着交换。原本按行存储的数据,瞬间就被解释成了按列。这个操作是零拷贝的,代价几乎为零。

import torch

t = torch.arange(6).reshape(2, 3)
print(t.shape)          # torch.Size([2, 3])
print(t.stride())       # (3, 1):沿第0维走一步跨3个元素,沿第1维走一步跨1个
print(t.storage_offset())  # 0:从Storage开头开始读

tt = t.t()
print(tt.shape)         # torch.Size([3, 2])
print(tt.stride())      # (1, 3):只是交换了步长,数据没有搬动
print(tt.untyped_storage().data_ptr() == t.untyped_storage().data_ptr())  # True,同一块存储

除了转置,PyTorch还提供了更底层的as_strided接口,允许直接指定shape、stride和offset来构造张量视图。卷积的im2col优化、滑动窗口提取等高性能技巧,底层都离不开它。不过要特别注意,如果给出的stride超出了Storage的实际范围,读取到的就是未定义的内存,程序可能不报错但结果完全错误,这是使用as_strided时最大的坑。

内存共享:视图操作的底层逻辑与风险

切片、view、reshape(在满足条件时)、expand、from_numpy这些操作有一个共同点:它们返回的新张量与原张量共享同一块Storage。这就是文章开头提到的现象的根源——修改新张量的值,就是直接改写共享的内存,原张量自然跟着变化。

import torch

t = torch.tensor([1.0, 2.0, 3.0, 4.0])

# 切片产生视图,共享底层存储
sub = t[1:3]
sub[0] = 99.0
print(t)   # tensor([ 1., 99.,  3.,  4.]),原张量被修改

# view同样共享存储
v = t.view(2, 2)
v[0, 0] = -1.0
print(t)   # tensor([-1., 99.,  3.,  4.])

# from_numpy与numpy()在CPU上也是零拷贝共享
import numpy as np
arr = np.array([1.0, 2.0])
ta = torch.from_numpy(arr)
arr[0] = 100.0
print(ta)  # tensor([100., 2.])

如果确实需要一个独立副本,应该使用clone()。clone会分配一块全新的Storage并把数据完整复制过去,新张量与原张量从此互不影响。另一个相关函数是copy_(),它是一个原地操作,把别的张量的值写进当前张量的Storage中,常用于把GPU数据搬回CPU或反向操作。简单记法:clone是给自己造一份新的,copy_是把别人的抄进自己家里。

这里还有一个容易混淆的点:reshape和view的区别。view严格要求张量在内存中是连续的(即stride满足行主序规则),否则直接报错;reshape则更宽容,能返回视图时就返回视图,不能时就会静默地复制一份数据再变形。所以用reshape得到的张量,有时共享存储有时不共享,行为不确定。如果代码逻辑依赖共享或依赖独立,一定要显式选择view或clone,避免埋下隐患。

序列化时Storage是如何被保存的

torch.save保存张量时,实际保存的是张量引用的那块Storage,而不是只保存张量可见范围内的元素。也就是说,如果你从一个大张量上切了一小片出来保存,落盘的可能是整个大Storage。这在保存模型中间结果或者缓存数据集切片时,会造成文件体积远超预期。

import torch

big = torch.arange(10000, dtype=torch.float32)
small = big[9000:]  # 只切出最后1000个元素

torch.save(small, "small.pt")
# 注意:保存的文件包含整个10000个元素的Storage,
# 而不是只有可见的1000个元素,文件比预期大约10倍

# 正确做法:先clone切断与大存储的关联
torch.save(small.clone(), "small_fixed.pt")

反过来,这个特性也有好处:如果多个张量共享同一块Storage,torch.save只会把这块Storage写入文件一次,加载回来后这些张量之间的共享关系会被完整还原。这在保存带有复杂视图关系的结构时非常省心。另外,对于特别大的数据,保存时可以指定mmap=True,加载时张量会通过内存映射方式访问磁盘文件,不必把全部内容一次性读进内存,适合处理超大规模的checkpoint或数据集。

常见问题解答

问题一:为什么修改切片的值,原张量也变了? 因为切片返回的是视图,新旧张量指向同一块Storage。解决方式是对切片调用clone(),得到真正独立的数据副本。

问题二:大张量切片后只留一小部分,为什么内存没有释放? 只要还有任何一个张量引用着那块Storage,整块内存就无法被回收。切片越小,这种浪费越明显。处理办法同样是clone()后丢弃原张量,或者用del及时解除所有引用。

问题三:怎么判断两个张量是否共享存储? 最直接的方式是比较底层存储的数据指针:a.data_ptr() == b.data_ptr()。也可以用a.untyped_storage().data_ptr() == b.untyped_storage().data_ptr(),后者判断的是是否指向同一块Storage,前者更严格,连偏移位置都要求一致。

问题四:调用tensor.storage()报错或警告怎么办? PyTorch 2.0之后,类型化的Storage API被逐步弃用,官方推荐改用tensor.untyped_storage()获取底层存储对象。旧代码升级时把这个调用替换掉即可,其余逻辑基本不用动。

问题五:怎么判断张量在内存中是否连续? 调用tensor.is_contiguous()即可。转置后的张量通常不连续,此时对其做view会报错,先调用contiguous()会复制出一份连续的副本,之后就能正常view了。理解这一点,能解决实际开发中相当一部分莫名其妙的报错。

总结一下,PyTorch的存储设计核心就是数据与描述分离:Storage负责装数据,Tensor负责解释数据。掌握shape、stride、storage_offset这三个元数据如何协作,再记住哪些操作共享存储、哪些操作复制数据,前面提到的那些奇怪现象就都能自己推导出来了。

PyTorch Storage张量存储内存共享修改时间:2026-09-26 02:43:27

免责声明:已尽一切努力确保本网站所含信息的准确性。网站作品多为原创整理与精心创作,观点力求客观中立。本站旨在免费分享,内容仅供个人学习、研究或参考使用。若引用了第三方作品,版权归原作者所有。如内容涉及您的权益,请联系我们进行处理Email:chomcom@qq.com。
引用或转载本作品时,请注明当前出处:https://www.ipipp.com/html/0926/61956.html,基于非商业用途的前提下,欢迎转载或二创本作品。
内容垂直聚焦
专注技术核心技术栏目,确保每篇文章深度聚焦于实用技能。从代码技巧到架构设计,为用户提供无干扰的纯技术知识沉淀,精准满足专业提升需求。
知识结构清晰
覆盖从开发到部署的全链路。AI、前端、编程、数据库、服务器、建站、系统层层递进,构建清晰学习路径,帮助用户系统化掌握开发与运维所需的核心技术。
深度技术解析
拒绝泛泛而谈,深入技术细节与实践难点。无论是数据库优化还是服务器配置,均结合真实场景与代码示例进行剖析,致力于提供可直接应用于工作的解决方案。
专业领域覆盖
精准对应开发生命周期。从前端界面到后端编程,从数据库操作到服务器运维,形成完整闭环,一站式满足全栈工程师和运维人员的技术需求。
即学即用高效
内容强调实操性,步骤清晰、代码完整。用户可根据教程直接复现和应用于自身项目,显著缩短从学习到实践的距离,快速解决开发中的具体问题。
持续更新保障
专注既定技术方向进行长期、稳定的内容输出。确保各栏目技术文章持续更新迭代,紧跟主流技术发展趋势,为用户提供经久不衰的学习价值。