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

如何不占用额外内存合并变量?解决torch.cat内存溢出问题

解决PyTorch合并大张量内存溢出问题

问题根源

torch.cat会创建一个全新的张量存储合并结果,这意味着合并过程中内存需要同时容纳两个3GB的源张量和一个6GB的目标张量,峰值占用达到12GB,远超8GB内存限制,必然导致溢出。

可行解决方案

方案1:预分配内存+分批加载(控制内存峰值)

先预分配好合并后的张量,加载第一个张量后直接复制到目标内存,释放原张量后再分块加载第二个张量并复制,全程内存峰值控制在7GB左右(6GB目标张量+1GB分块)。

import torch

# 获取两个特征文件的形状信息(提前确认样本数,避免重复加载全量数据)
temp_feat1 = torch.load('feat1.pt')
feat1_shape = temp_feat1.shape
del temp_feat1
torch.cuda.empty_cache()  # CPU环境可省略,加速内存回收

temp_feat2 = torch.load('feat2.pt')
feat2_shape = temp_feat2.shape
del temp_feat2
torch.cuda.empty_cache()

# 预分配合并后的张量,与源张量同类型、设备
merged = torch.empty(
    (feat1_shape[0] + feat2_shape[0],) + feat1_shape[1:],
    dtype=temp_feat1.dtype,
    device=temp_feat1.device
)

# 加载第一个特征并复制到目标张量
feat1 = torch.load('feat1.pt')
merged[:feat1_shape[0]] = feat1
del feat1
torch.cuda.empty_cache()

# 分块加载第二个特征并复制,调整chunk_size控制单块内存占用
chunk_size = 30000  # 根据你的张量维度调整,确保单块内存≤1GB
start_pos = feat1_shape[0]
for i in range(0, feat2_shape[0], chunk_size):
    end = min(i + chunk_size, feat2_shape[0])
    feat2_chunk = torch.load('feat2.pt')[i:end]
    merged[start_pos + i:start_pos + end] = feat2_chunk
    del feat2_chunk
    torch.cuda.empty_cache()

方案2:内存映射加载(按需读取磁盘数据)

利用PyTorch的内存映射功能,将特征文件以只读映射方式加载,数据实际存储在磁盘,仅当访问特定块时才加载到内存,合并时全程内存占用仅为目标张量大小+当前分块大小。

import torch

# 内存映射加载两个特征文件(仅加载元数据,数据留在磁盘)
feat1_mmap = torch.load('feat1.pt', map_location='cpu', mmap=True)
feat2_mmap = torch.load('feat2.pt', map_location='cpu', mmap=True)

# 预分配合并后的张量
merged = torch.empty(
    (feat1_mmap.shape[0] + feat2_mmap.shape[0],) + feat1_mmap.shape[1:],
    dtype=feat1_mmap.dtype
)

# 分块复制第一个特征
chunk_size = 50000
for i in range(0, feat1_mmap.shape[0], chunk_size):
    end = min(i + chunk_size, feat1_mmap.shape[0])
    merged[i:end] = feat1_mmap[i:end]

# 分块复制第二个特征
start_pos = feat1_mmap.shape[0]
for i in range(0, feat2_mmap.shape[0], chunk_size):
    end = min(i + chunk_size, feat2_mmap.shape[0])
    merged[start_pos + i:start_pos + end] = feat2_mmap[i:end]

# 释放内存映射对象
del feat1_mmap, feat2_mmap

注意事项

  • 如果使用GPU,确保map_location指定正确设备,且torch.cuda.empty_cache()要及时调用,释放无用显存。
  • 调整chunk_size时,要根据你的张量具体维度计算单块内存占用,确保峰值不超过8GB。

内容的提问来源于stack exchange,提问作者YA xiang

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.11 10:42:42