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

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内存的方法

  1. 彻底清理张量引用
    仅删除列表不足以释放内存,需确保单批张量的所有引用都被清除。处理时先拆解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()
    
  2. 避免隐式设备拷贝
    PyTorch Lightning的trainer.predict()默认会将GPU张量转移到CPU,可能产生隐式拷贝残留。可以在模型的predict_step中主动控制输出处理,或者在获取结果后立即转换为numpy数组(如果后续不需要张量操作),减少内存占用。

二、更内存高效的大规模推理方案

  1. 边推理边写入磁盘
    无需将所有批次结果存放在内存中,每处理完一批就把结果转成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()
    
  2. 裁剪不必要的输出
    如果只需要部分隐藏层的输出,修改模型forward方法只返回目标层,减少单批结果的内存占用:
    def forward(self, x):
        outputs = self.model(x, output_hidden_states=True)
        # 只返回第10到12层的隐藏状态
        return (outputs.last_hidden_state, outputs.hidden_states[-3:])
    
  3. 启用半精度推理
    使用FP16半精度推理可大幅降低GPU和CPU的内存占用,PyTorch Lightning只需在初始化Trainer时添加参数:
    trainer = pl.Trainer(accelerator="gpu", precision=16)
    
  4. 迭代式拼接结果
    每处理一批就直接和已有的结果拼接,然后立即删除当前批次的张量,避免保存所有批次再统一拼接:
    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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.25 04:13:25