You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何以低内存开销在指定索引处高效合并两个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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.29 02:42:28