PyTorch+Colab环境下深度学习模型微调内存管理咨询
深度学习模型微调内存优化方案(针对Colab PyTorch OOM问题)
问题背景
在Google Colab环境中使用PyTorch微调DistilGPT2模型时,出现RAM耗尽导致运行时会话自动关闭的问题,切换至GPU后问题仍未解决。运行代码如下:
!pip install datasets !pip install transformers import torch import wandb import torch.nn as nn from datasets import load_dataset from transformers import TrainingArguments, Trainer, DataCollatorForLanguageModeling from transformers import AutoTokenizer, AutoModelForCausalLM, AutoConfig !huggingface-cli login model_name = "distilgpt2" print(f"Using GPU: {torch.cuda.is_available()}") tokenizer = AutoTokenizer.from_pretrained(model_name) len(tokenizer) tokenizer.add_special_tokens({'pad_token': '[PAD]'}) len(tokenizer) class Custom_GPT2_Model(nn.Module): def __init__(self, tokenizer): super().__init__() self.gpt2 = AutoModelForCausalLM.from_pretrained(model_name) self.gpt2.resize_token_embeddings(len(tokenizer)) for param in self.gpt2.parameters(): param.requires_grad = False self.gpt2.gradient_checkpointing_enable() self.custom_layer = nn.Linear(self.gpt2.config.vocab_size, self.gpt2.config.vocab_size) def forward(self, input_ids, attention_mask=None, labels=None): outputs = self.gpt2(input_ids=input_ids, attention_mask=attention_mask) logits = self.custom_layer(outputs.logits) loss = None if labels is not None: loss_func = nn.CrossEntropyLoss() # loss = loss_func(logits, labels) loss = loss_func(logits.view(-1, logits.size(-1)), labels.view(-1)) return { 'loss': loss, 'logits': logits } model = Custom_GPT2_Model(tokenizer) # Data Pre-Processing dataset = load_dataset("wikitext", "wikitext-2-raw-v1") dataset def tokenize_func(examples): return tokenizer(examples['text'], padding='max_length', truncation=True, max_length=128, return_tensors='pt') tokenized_dataset = dataset.map(tokenize_func, batched=True, remove_columns=['text']) tokenized_dataset training_args = TrainingArguments( output_dir='./results', num_train_epochs=0.5, per_device_train_batch_size=1, gradient_accumulation_steps=16, fp16=True, warmup_steps=2500, learning_rate=0.1, weight_decay=0.01, logging_dir='./logs', logging_steps=5000, ) trainer = Trainer( model = model, args = training_args, train_dataset = tokenized_dataset['train'].select(range(100)), eval_dataset = tokenized_dataset['test'].select(range(100)), data_collator = DataCollatorForLanguageModeling(tokenizer=tokenizer, mlm=False) ) trainer.train()
内存优化方案
1. 数据预处理优化
当前预处理存在两个核心内存浪费点:padding='max_length'产生大量冗余填充张量,return_tensors='pt'直接将数据集转为PyTorch张量占用内存。修改后代码:
def tokenize_func(examples): # 按批次内最长样本填充,减少冗余 return tokenizer(examples['text'], padding='longest', truncation=True, max_length=128) # 控制map批次大小,禁用缓存减少内存占用 tokenized_dataset = dataset.map(tokenize_func, batched=True, batch_size=1000, remove_columns=['text'], load_from_cache_file=False) # 清理原数据集释放内存 del dataset import gc gc.collect()
2. 训练参数与模型配置优化
- 修正不合理的训练参数:
warmup_steps远大于实际训练步数,过大学习率既影响收敛也增加内存开销; - 强制模型移至GPU,避免跨设备张量拷贝;
- 关闭不必要的评估与日志,减少内存占用。
修改后代码:
# 将模型移至GPU model = model.to('cuda') if torch.cuda.is_available() else model training_args = TrainingArguments( output_dir='./results', num_train_epochs=0.5, per_device_train_batch_size=1, gradient_accumulation_steps=16, fp16=True, warmup_steps=5, # 匹配实际训练步数 learning_rate=5e-5, # 因果语言模型微调合理学习率 weight_decay=0.01, logging_dir='./logs', logging_steps=1, eval_strategy='no', # 关闭评估节省内存 load_best_model_at_end=False, report_to='none' # 关闭wandb日志,避免额外内存消耗 )
3. 显存与内存主动清理
在训练关键节点主动清理无用张量与显存,避免内存泄漏:
# 训练前清理显存 if torch.cuda.is_available(): torch.cuda.empty_cache() trainer.train() # 训练后清理显存 if torch.cuda.is_available(): torch.cuda.empty_cache()
4. 模型结构微调
将损失函数实例化移至__init__中,减少重复实例化的内存开销:
class Custom_GPT2_Model(nn.Module): def __init__(self, tokenizer): super().__init__() self.gpt2 = AutoModelForCausalLM.from_pretrained(model_name) self.gpt2.resize_token_embeddings(len(tokenizer)) for param in self.gpt2.parameters(): param.requires_grad = False self.gpt2.gradient_checkpointing_enable() self.custom_layer = nn.Linear(self.gpt2.config.vocab_size, self.gpt2.config.vocab_size) self.loss_func = nn.CrossEntropyLoss() # 移至初始化阶段 def forward(self, input_ids, attention_mask=None, labels=None): outputs = self.gpt2(input_ids=input_ids, attention_mask=attention_mask) logits = self.custom_layer(outputs.logits) loss = None if labels is not None: loss = self.loss_func(logits.view(-1, logits.size(-1)), labels.view(-1)) return {'loss': loss, 'logits': logits}
内容的提问来源于stack exchange,提问作者AaKash Singh
相关产品推荐
相关产品推荐

