无需复制交换PyTorch张量数据的正确方法及相关疑问
PyTorch张量底层数据无复制交换方案
正确实现方式
要无复制交换两个同维度PyTorch张量x、y的底层数据,核心是交换它们对底层Storage对象的绑定关系。推荐使用PyTorch公开的set_()方法,避免依赖私有属性导致版本兼容问题:
# 保存x的原始存储信息 x_storage = x.storage() x_offset = x.storage_offset() x_size = x.size() # 将x重新绑定到y的存储 x.set_(y.storage(), y.storage_offset(), y.size()) # 将y重新绑定到x的原始存储 y.set_(x_storage, x_offset, x_size)
如果图省事,也可以直接交换两个Storage的底层数据指针(注意这是访问私有属性,不同PyTorch版本可能有变动):
x_storage = x.storage() y_storage = y.storage() x_storage._data_ptr, y_storage._data_ptr = y_storage._data_ptr, x_storage._data_ptr
两种方式都不会复制数据,只是修改了张量与底层存储的映射关系。
对持有张量引用的代码的影响
完全符合预期,不会破坏现有代码。所有持有x引用的代码,后续对x的读写操作都会直接作用于原来y的底层数据;持有y引用的代码则会操作原来x的数据。因为这些代码持有的是张量对象的引用,而张量对象的底层存储绑定关系已经被交换,无需修改这些代码就能实现需求。
引用计数与GC的工作状态
交换操作不会影响PyTorch的引用计数和垃圾回收机制。Storage对象的引用计数会随着张量的绑定关系变化自动更新:当某个Storage不再被任何张量(包括x、y及其他引用)引用时,Python的GC会正常回收对应的内存资源,不会出现内存泄漏或引用计数混乱的问题。
内容的提问来源于stack exchange,提问作者AaronDefazio
相关产品推荐
相关产品推荐

