PyTorch使用索引数组选取Tensor切片无法深拷贝的解决方案咨询
PyTorch 索引拷贝问题解决方案
首先纠正你测试里的结论偏差:你对两种索引的内存共享特性判断完全反了。
PyTorch的索引内存规则非常明确:
- 用
[start:end:step]语法做的基础切片,返回的是原Tensor的视图,和原Tensor共享底层内存,修改切片值会同步改动原Tensor - 用索引列表/布尔数组做的高级索引(也就是你写的
a[[2,3]]),原生返回的就是和原Tensor内存隔离的新Tensor,默认就不会共享存储。
你可以用下面的代码复现验证:
import torch a = torch.Tensor([1,2,3,4,5,6]) # 测试普通切片 b = a[2:4] b[0] = 999 print(a) # 输出tensor([ 1., 2., 999., 4., 5., 6.]),证明b和a共享内存 # 重置a,测试索引数组选取 a = torch.Tensor([1,2,3,4,5,6]) c = a[[2,3]] c[0] = 999 print(a) # 输出tensor([1., 2., 3., 4., 5., 6.]),证明c和a内存独立
如果你在实际复杂业务逻辑里,遇到索引数组选取结果仍和原Tensor共享内存的特殊情况(一般是旧版本PyTorch逻辑差异、或者索引操作嵌套在视图转换链路里导致的),要强制拿到完全独立的深拷贝结果,直接调用PyTorch原生的.clone()方法即可,这是做Tensor深拷贝的标准API:
# 无论前面的索引返回的是视图还是拷贝,.clone()都会生成底层存储完全独立的新Tensor c = a[[2,3]].clone()
你之前尝试的.reshape()、.view()、.contiguous()都无法保证拿到独立拷贝,原因很简单:
.view()仅做张量形状逻辑转换,永远不会复制底层数据,只在内存连续时可用.reshape()优先返回视图,仅在无法生成合法视图时才会隐式复制数据,结果是否独立完全不可控.contiguous()仅在原张量内存不连续时才会复制数据,内存连续时直接返回原张量的视图,同样不做数据拷贝
补充:如果你不需要保留张量对应的自动微分计算图链路,可以在
.clone()后追加.detach(),即a[[2,3]].clone().detach(),得到和原计算图、原存储完全隔离的纯数值张量。
内容的提问来源于stack exchange,提问作者Yandong Ji
相关产品推荐
相关产品推荐

