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时几乎立即溢出,但在更小数据集上该批量大小可正常运行
需求:排查问题原因、计算单数据点内存占用、寻求更高效运行方案
问题原因排查
- 重复前向传播的显存泄漏
代码对每个batch执行16次重复前向传播,PyTorch默认会保留前向传播的计算图用于梯度回溯,即使未显式计算梯度。随着循环次数增加,中间张量持续堆积,显存被逐步占满,这就是batch size=32时前几轮正常、后续才报错的核心原因。 - 数据集规模的线性影响
小数据集下,即使batch size=64,总数据量小,中间张量堆积速度慢,不会立刻触发溢出;但数据集变大后,每轮循环产生的中间张量总量超过显存阈值,就会立即报错。 empty_cache()的局限性torch.cuda.empty_cache()仅释放CUDA标记为已用但未被PyTorch引用的显存,无法清理PyTorch仍在追踪的中间张量,只能临时缓解显存压力,且频繁调用会增加CUDA上下文切换开销,拖慢运行速度。
单数据点内存占用计算
单数据点的显存占用由三部分组成:
- 模型参数内存:固定值,比如BERT-base约110M参数,FP32存储时为
110M * 4字节 = 440MB,FP16存储时减半为220MB。 - 输入张量内存:
input_ids和attention_mask均为长度为max_seq_len的整数张量(int64,8字节/元素)。以max_seq_len=512为例,单个数据点输入张量内存为512*8*2 = 8192字节 ≈ 8KB。 - 中间计算张量内存:前向传播时模型内部生成的隐藏层、注意力权重等(即使设置
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
相关产品推荐
相关产品推荐

