PyTorch大规模推理中CPU内存释放方法及更优方案咨询
解决BERT大规模推理的CPU内存不足(OOM)问题
问题描述
我在单台机器上对预训练BERT模型执行大规模推理时,遭遇了CPU内存不足(OOM)错误。由于数据集过大无法一次性完成推理,我采用分批处理的方式,将每批结果存入列表,最后再拼接这些张量。我清楚将张量存储在列表中会快速占用大量CPU内存,但拼接后找不到释放这些内存的有效方法,导致后续流程仍会触发OOM错误。
最小复现代码
import gc, time, torch, pytorch_lightning as pl from transformers import BertTokenizer, BertModel from torch.utils.data import DataLoader class EncoderModelPL(pl.LightningModule): def __init__( self, model: BertModel, ): super(EncoderModelPL, self).__init__() self.model: BertModel = model def forward(self, x): return self.model(x, output_hidden_states=True) # 下游任务需要中间隐藏层,因此必须返回这些状态 MODEL_ID = "bert-base-uncased" tokenizer = BertTokenizer.from_pretrained(MODEL_ID) model = EncoderModelPL(BertModel.from_pretrained(MODEL_ID)).to("cuda") dataset = torch.randint(low=0, high=30000, size=(800, 200), device="cuda") tokens_dataloader = DataLoader(dataset, batch_size=32, shuffle=False) trainer = pl.Trainer(accelerator="gpu") bert_outputs_per_batch: list = trainer.predict( model=model, dataloaders=tokens_dataloader ) # CPU内存持续增长,这里输出一个长度等于批次数的列表,每个元素是单批BERT输出,存储在CPU上 del bert_outputs_per_batch gc.collect() ...<后续流程>...
已尝试的方法
- 删除结果列表并执行垃圾回收(如代码所示)
- 操作后休眠几秒辅助垃圾回收,无效果
- 给dataset添加
.detach()确认不是计算图追踪问题,问题依旧存在 - 在DataLoader中设置
pin_memory=False,无明显变化
核心问题
- 如何强制释放CPU内存中的张量?
- 有没有比分批处理更内存高效的大规模推理方案?
解决方案
一、强制释放CPU内存的方法
- 彻底清理张量引用
仅删除列表不足以释放内存,需确保单批张量的所有引用都被清除。处理时先拆解BERT输出的结构(比如hidden_states是张量元组),逐个删除张量后再清理列表:bert_outputs_per_batch = trainer.predict(model=model, dataloaders=tokens_dataloader) # 逐个清理单批输出中的张量 for batch_out in bert_outputs_per_batch: # 遍历所有隐藏层张量并删除 for tensor in batch_out.hidden_states: del tensor del batch_out # 删除列表本身 del bert_outputs_per_batch # 触发Torch和系统的垃圾回收 torch.cuda.empty_cache() gc.collect() - 避免隐式设备拷贝
PyTorch Lightning的trainer.predict()默认会将GPU张量转移到CPU,可能产生隐式拷贝残留。可以在模型的predict_step中主动控制输出处理,或者在获取结果后立即转换为numpy数组(如果后续不需要张量操作),减少内存占用。
二、更内存高效的大规模推理方案
- 边推理边写入磁盘
无需将所有批次结果存放在内存中,每处理完一批就把结果转成numpy数组写入磁盘(如HDF5、NPY格式),后续需要时再从磁盘读取拼接:import h5py with h5py.File("bert_hidden_states.h5", "w") as f: for batch_idx, batch_out in enumerate(trainer.predict(model=model, dataloaders=tokens_dataloader)): # 按层写入每个批次的隐藏状态 for layer_idx, hidden_state in enumerate(batch_out.hidden_states): ds_name = f"layer_{layer_idx}/batch_{batch_idx}" f.create_dataset(ds_name, data=hidden_state.detach().cpu().numpy()) # 立即释放当前批次内存 del batch_out gc.collect() - 裁剪不必要的输出
如果只需要部分隐藏层的输出,修改模型forward方法只返回目标层,减少单批结果的内存占用:def forward(self, x): outputs = self.model(x, output_hidden_states=True) # 只返回第10到12层的隐藏状态 return (outputs.last_hidden_state, outputs.hidden_states[-3:]) - 启用半精度推理
使用FP16半精度推理可大幅降低GPU和CPU的内存占用,PyTorch Lightning只需在初始化Trainer时添加参数:trainer = pl.Trainer(accelerator="gpu", precision=16) - 迭代式拼接结果
每处理一批就直接和已有的结果拼接,然后立即删除当前批次的张量,避免保存所有批次再统一拼接:final_hidden_states = None for batch_out in trainer.predict(model=model, dataloaders=tokens_dataloader): current_hidden = batch_out.hidden_states if final_hidden_states is None: # 初始化结果列表 final_hidden_states = [h.detach().cpu() for h in current_hidden] else: # 逐层拼接 for layer_idx in range(len(final_hidden_states)): final_hidden_states[layer_idx] = torch.cat( [final_hidden_states[layer_idx], current_hidden[layer_idx].detach().cpu()], dim=0 ) # 清理当前批次内存 del batch_out gc.collect()
内容的提问来源于stack exchange,提问作者Kevin
相关产品推荐
相关产品推荐

