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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.15 11:37:31