PyTorch中非单例维度扩展如何避免内存数据拷贝?
非连续张量维度扩展无内存拷贝解决方案
针对形状为[a,b,c]的非连续张量s(b远大于1),要在第二维度重复n次得到[a,nb,c]的张量且避免内存拷贝,最优方案是直接利用torch.Tensor.as_strided构造视图,全程无内存拷贝。
具体实现代码
import torch # 示例参数:替换成你实际的a、b、c、n值 a, b, c, n = 2, 1000, 3, 5 # 构造一个非连续张量示例(模拟你的输入张量s) s = torch.randn(a, b, c)[..., ::2].transpose(1, 2) # 使用as_strided创建无拷贝的重复视图 new_shape = (a, b * n, c) # 步长沿用原张量的步长:第二维度重复时,步长不变,仅形状扩展 strides = s.stride() result = s.as_strided(new_shape, strides)
方案说明
as_strided仅修改张量的元信息(形状、步长),完全复用原张量的内存块,没有任何内存拷贝操作,速度拉满。- 合法性验证:可以通过
result.storage().data_ptr() == s.storage().data_ptr()确认两者共享同一块内存。 - 注意事项:必须保证步长和形状的设置合法,避免越界访问内存。这里因为是在第二维度重复,直接沿用原步长即可,每个元素会被连续访问
n次,最终呈现出重复的效果。
现有方法产生拷贝的原因
repeat_interleave:本质是对元素进行物理复制,必须开辟新内存存储复制后的内容,必然产生拷贝。expand后用view/reshape:expand本身是视图,但reshape或view处理非连续张量时,会自动触发contiguous()操作,而contiguous()会将张量重新排列为连续内存布局,这一步就会产生内存拷贝。
内容的提问来源于stack exchange,提问作者G.G.
相关产品推荐
相关产品推荐

