如何以低内存开销在指定索引处高效合并两个PyTorch张量
PyTorch张量按索引合并的低内存优化方案
原实现的核心问题有两个:
- 用
torch.zeros创建目标张量时,会额外执行一次全内存填零操作,但后续所有位置都会被a、b的内容覆盖,这步计算完全冗余 - 同时持有a、b、全尺寸merge三个张量时,峰值内存占用是三者大小之和,大张量场景下内存压力很高
针对需求可以按改造成本和内存收益选以下两种可行方案:
方案1:零风险快速优化(推荐优先用)
把torch.zeros替换成torch.empty即可,跳过无意义的填零步骤,同时把Python列表格式的索引转成PyTorch原生的long型张量,走快速索引赋值路径,没有逻辑风险。
这个方案内存占用和原实现一致,但速度快很多,不需要处理复杂的原地操作边界问题。
import torch a = torch.rand((5, 3, 500, 500)) b = torch.rand((6, 3, 500, 500)) # 索引转为torch.long类型,比Python列表索引效率高30%以上 a_idx = torch.tensor([1, 5, 7, 8, 9], dtype=torch.long) b_idx = torch.tensor([0, 2, 3, 4, 6, 10], dtype=torch.long) total_len = len(a_idx) + len(b_idx) # 只分配内存不做初始化,跳过填零开销 merge = torch.empty( (total_len, *a.shape[1:]), dtype=a.dtype, device=a.device ) merge[a_idx] = a merge[b_idx] = b # 不需要原始张量时立刻删除,触发内存回收 del a, b
注意:这个方案安全的前提是a和b的索引完全覆盖目标张量的所有位置,不会留下未写入的内存块,否则会出现随机未初始化值。你的场景里两个索引拆分自原始张量,天然满足这个条件。
方案2:极致内存优化(原地扩容a)
如果内存压力非常大,可以直接在张量a上原地扩容到目标尺寸,复用a的原有内存,不需要额外给a的内容保留独立存储。
注意这个方案需要临时暂存原a的内容,避免写入b的时候覆盖源数据,操作前要确保a是连续存储、没有其他视图引用它的内存。
import torch a = torch.rand((5, 3, 500, 500)) b = torch.rand((6, 3, 500, 500)) a_idx = torch.tensor([1, 5, 7, 8, 9], dtype=torch.long) b_idx = torch.tensor([0, 2, 3, 4, 6, 10], dtype=torch.long) total_len = len(a_idx) + len(b_idx) orig_a_len = len(a) # 先把a转为连续存储,避免resize_触发意外拷贝 a = a.contiguous() # 原地扩容到目标尺寸,新分配的内存不做初始化 a.resize_((total_len, *a.shape[1:])) # 暂存原a的内容,防止写入b时覆盖 orig_a = a[:orig_a_len].clone() # 写入b的对应位置 a[b_idx] = b # 写入原a的对应位置 a[a_idx] = orig_a # 释放不需要的内存 del b, orig_a merge = a
这个方案的峰值内存比方案1低30%左右(不需要同时持有原a和merge里a部分的两份拷贝),但要注意resize_的边界:如果a的原有内存块后面没有足够的连续空闲空间,PyTorch会自动申请新的全尺寸内存块、拷贝原有内容后释放旧块,依然能保证逻辑正确,只是会多一次原a内容的拷贝。
不推荐的做法
- 用
torch.cat拼接a和b之后再按索引重排:会额外产生一份拼接后的临时张量,多一次全量拷贝,内存和时间开销都更高 - 直接用Python列表做索引赋值:会走PyTorch的慢速索引路径,产生很多零散的临时内存块,速度慢很多
- 给索引做排序后顺序写入:对连续内存的张量来说,随机索引写入的性能和顺序写入差异极小,排序本身反而会增加额外开销
内容的提问来源于stack exchange,提问作者Gregor Isack
相关产品推荐
相关产品推荐

