如何不占用额外内存合并变量?解决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
相关产品推荐
相关产品推荐

