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

PyTorch Hugging Face评估阶段显存耗尽问题求助

评估阶段CUDA显存耗尽问题排查与解决

使用80GB显存GPU训练时epoch运行正常,但评估阶段(训练集与验证集规模大致相同)出现显存耗尽,报错信息如下:

File "/home.../transformers/trainer_pt_utils.py", line 75, in torch_pad_and_concatenate
return torch.cat((tensor1, tensor2), dim=0)
RuntimeError: CUDA out of memory. Tried to allocate 33.84 GiB (GPU 0; 79.35 GiB total 
capacity; 36.51 GiB already allocated; 32.48 GiB free; 44.82 GiB reserved in total by 
PyTorch) If reserved memory is >> allocated memory try setting max_split_size_mb to 
avoid fragmentation.  See documentation for Memory Management and 
PYTORCH_CUDA_ALLOC_CONF

训练与验证数据创建代码

train_texts, train_labels = read_dataset('basic_train.tsv') 

val_texts, val_labels = read_dataset('basic_val.tsv')  

train_encodings = tokenizer(train_texts, truncation=False, padding=True) 
val_encodings = tokenizer(val_texts, truncation=False, padding=True)

class Dataset(torch.utils.data.Dataset):     
    def __init__(self, encodings, labels):         
        self.encodings = encodings         
        self.labels = labels 
         ...         
        return item 

train_dataset = Dataset(train_encodings, train_labels) 
val_dataset = Dataset(val_encodings, val_labels) 

训练代码

training_args = TrainingArguments(
output_dir='./results',          
num_train_epochs=10,             
per_device_train_batch_size=8,  
per_device_eval_batch_size=8,   
warmup_steps=500,                
weight_decay= 5e-5,              
logging_dir='./logs',            
logging_steps=10,
learning_rate= 2e-5,
eval_steps= 100,
save_steps=30000,
evaluation_strategy= 'steps'
)
model = AutoModelForSeq2SeqLM.from_pretrained("t5-base")


metric = load_metric('accuracy')

def compute_metrics(eval_pred):
  predictions, labels = eval_pred
  predictions = np.argmax(predictions, axis=1)
  return metric.compute(predictions=predictions, references=labels)

def collate_fn_t5(batch):
  input_ids = torch.stack([example['input_ids'] for example in batch])
  attention_mask = torch.stack([example['attention_mask'] for example in batch])
  labels = torch.stack([example['input_ids'] for example in batch])
   return {'input_ids': input_ids, 'attention_mask': attention_mask, 'labels': labels}


trainer = Trainer(
model=model,                       
args=training_args,                  
train_dataset=train_dataset,         
eval_dataset=val_dataset,
compute_metrics=compute_metrics,
data_collator=collate_fn_t5,
        # evaluation dataset
 )

trainer.train()

eval_results = trainer.evaluate()

解决方案

  • 修正数据预处理的padding策略
    当前truncation=False, padding=True会将整个数据集padding到样本的最大长度,若验证集存在超长样本,会导致单batch显存占用剧增。改为按固定最大长度截断并padding,匹配T5模型的默认输入限制:

    train_encodings = tokenizer(train_texts, truncation=True, padding="max_length", max_length=512)
    val_encodings = tokenizer(val_texts, truncation=True, padding="max_length", max_length=512)
    
  • 降低评估批次大小
    评估阶段无需保持和训练相同的批次大小,将per_device_eval_batch_size从8下调至4或更小,直接减少单batch的显存占用:

    per_device_eval_batch_size=4
    
  • 修复collate_fn的逻辑错误
    你的collate_fn_t5中错误地将labels设置为input_ids,这不仅导致模型训练逻辑错误,还会引入冗余tensor占用额外显存,应改为使用样本真实标签:

    def collate_fn_t5(batch):
      input_ids = torch.stack([example['input_ids'] for example in batch])
      attention_mask = torch.stack([example['attention_mask'] for example in batch])
      labels = torch.stack([example['labels'] for example in batch])  # 修正为真实labels
      return {'input_ids': input_ids, 'attention_mask': attention_mask, 'labels': labels}
    
  • 优化显存分配缓解碎片
    根据报错提示,设置环境变量优化PyTorch的显存分配策略,减少碎片:

    export PYTORCH_CUDA_ALLOC_CONF=max_split_size_mb:128
    

    或在代码开头添加:

    import os
    os.environ['PYTORCH_CUDA_ALLOC_CONF'] = 'max_split_size_mb:128'
    
  • 控制评估预测的内存占用
    若验证集规模较大,Trainer默认保存所有预测结果会占用显存。对于分类任务,可显式关闭生成式预测,避免不必要的内存消耗:

    training_args = TrainingArguments(
        # ...其他参数
        predict_with_generate=False
    )
    

内容的提问来源于stack exchange,提问作者Chan Wing

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.27 01:02:51