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

Huggingface Transformers评估数据集规模与GPU显存溢出问题排查

问题描述

我基于Huggingface Transformers训练完成了BertForSequenceClassification模型,需要用它对不同数据执行大量前向传播。在优化批量大小时,出现异常GPU显存溢出情况,使用的是80G显存的A100 GPU,核心代码如下:

model = geneformer_utils.get_model(path=MODEL, output_attentions=False, enable_random_dropout=True)
training_args_dict = geneformer_utils.get_training_args(
    output_dir=OUTPUT,
    per_device_train_batch_size=1,
    per_device_eval_batch_size=16,
    eval_accumulation_steps=1,
)
training_args = transformers.training_args.TrainingArguments(**training_args_dict)
trainer = transformers.Trainer(
    model=model,
    args=training_args,
    data_collator=geneformer.DataCollatorForCellClassification(),
    train_dataset=test,
    eval_dataset=test,
    compute_metrics=None
)

data_loader = trainer.get_eval_dataloader()

full_result = []
for i in range(16):
    # 0: logits
    full_result.append([[]])
for _, batch in enumerate(data_loader):
    for i in range(16):
        gene_ids, attention_mask = batch['input_ids'], batch['attention_mask']
        model_predictions = model(
            input_ids=gene_ids,
            attention_mask=attention_mask,
            output_attentions=False,
            output_hidden_states=False
        )
        logits = list(model_predictions['logits'].cpu().detach().numpy())
        full_result[i][0].extend(logits)

异常表现:

  • 设per_device_eval_batch_size=32时,前4-5轮循环正常,之后才触发显存溢出
  • 插入torch.cuda.empty_cache()可避免或延迟报错,但会降低循环速度
  • 设per_device_eval_batch_size=64时几乎立即溢出,但在更小数据集上该批量大小可正常运行

需求:排查问题原因、计算单数据点内存占用、寻求更高效运行方案


问题原因排查

  1. 重复前向传播的显存泄漏
    代码对每个batch执行16次重复前向传播,PyTorch默认会保留前向传播的计算图用于梯度回溯,即使未显式计算梯度。随着循环次数增加,中间张量持续堆积,显存被逐步占满,这就是batch size=32时前几轮正常、后续才报错的核心原因。
  2. 数据集规模的线性影响
    小数据集下,即使batch size=64,总数据量小,中间张量堆积速度慢,不会立刻触发溢出;但数据集变大后,每轮循环产生的中间张量总量超过显存阈值,就会立即报错。
  3. empty_cache()的局限性
    torch.cuda.empty_cache()仅释放CUDA标记为已用但未被PyTorch引用的显存,无法清理PyTorch仍在追踪的中间张量,只能临时缓解显存压力,且频繁调用会增加CUDA上下文切换开销,拖慢运行速度。

单数据点内存占用计算

单数据点的显存占用由三部分组成:

  1. 模型参数内存:固定值,比如BERT-base约110M参数,FP32存储时为110M * 4字节 = 440MB,FP16存储时减半为220MB。
  2. 输入张量内存:input_ids和attention_mask均为长度为max_seq_len的整数张量(int64,8字节/元素)。以max_seq_len=512为例,单个数据点输入张量内存为512*8*2 = 8192字节 ≈ 8KB。
  3. 中间计算张量内存:前向传播时模型内部生成的隐藏层、注意力权重等(即使设置output_attentions=False和output_hidden_states=False,这些张量仍会用于计算)。这是显存占用的大头,BERT-base单数据点FP32下约占10-20MB,FP16下减半。

实际批量运行时,模型参数仅加载一次,显存增长主要来自输入和中间张量,总显存占用(FP32)≈模型参数内存 + (输入张量内存 + 中间张量内存)*batch size。


高效运行方案

1. 禁用计算图追踪

添加torch.no_grad()上下文管理器,阻止PyTorch生成计算图,彻底避免中间张量堆积,这是解决显存泄漏最有效的方法:

full_result = []
for i in range(16):
    full_result.append([[]])

# 关键:用torch.no_grad()包裹前向传播循环
with torch.no_grad():
    for _, batch in enumerate(data_loader):
        gene_ids, attention_mask = batch['input_ids'], batch['attention_mask']
        for i in range(16):
            model_predictions = model(
                input_ids=gene_ids,
                attention_mask=attention_mask,
                output_attentions=False,
                output_hidden_states=False
            )
            logits = model_predictions['logits'].cpu().detach().numpy().tolist()
            full_result[i][0].extend(logits)

2. 优化重复前向传播逻辑

如果16次前向传播是为了实现蒙特卡洛Dropout(因设置了enable_random_dropout=True),可将输入张量重复拼接后一次性前向传播,减少GPU调用开销:

import numpy as np

full_result = []
for i in range(16):
    full_result.append([[]])

with torch.no_grad():
    for _, batch in enumerate(data_loader):
        gene_ids, attention_mask = batch['input_ids'], batch['attention_mask']
        # 在batch维度重复输入张量16次
        repeated_input_ids = gene_ids.repeat(16, 1)
        repeated_attention_mask = attention_mask.repeat(16, 1)
        # 一次性完成16次前向传播
        model_predictions = model(
            input_ids=repeated_input_ids,
            attention_mask=repeated_attention_mask,
            output_attentions=False,
            output_hidden_states=False
        )
        logits = model_predictions['logits'].cpu().detach().numpy()
        # 拆分16组结果
        split_logits = np.array_split(logits, 16)
        for i in range(16):
            full_result[i][0].extend(split_logits[i].tolist())

3. 启用混合精度推理

将模型和输入转为FP16,可将显存占用降低约50%:

from torch.cuda.amp import autocast

# 模型参数转为FP16
model = model.half()

full_result = []
for i in range(16):
    full_result.append([[]])

with torch.no_grad(), autocast():
    for _, batch in enumerate(data_loader):
        gene_ids = batch['input_ids']
        # attention_mask转为FP16,input_ids为整数无需转换
        attention_mask = batch['attention_mask'].half()
        for i in range(16):
            model_predictions = model(
                input_ids=gene_ids,
                attention_mask=attention_mask,
                output_attentions=False,
                output_hidden_states=False
            )
            # 输出转float避免numpy精度问题
            logits = model_predictions['logits'].cpu().detach().float().numpy().tolist()
            full_result[i][0].extend(logits)

4. 调整数据加载器参数

  • 关闭pin_memory:大batch size下,pin_memory=True会占用额外CPU内存并增加GPU拷贝开销,可改为pin_memory=False
  • 优化num_workers:适当提高数据加载进程数(不超过CPU核心数),避免数据加载成为运行瓶颈

内容的提问来源于stack exchange,提问作者Nikolay Markov

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.16 19:16:01