如何为Flan T5-base模型配置DeepSpeed激活Checkpointing?
解决PyTorch Lightning + DeepSpeed下Flan-T5-base的显存不足问题
要在Flan-T5-base这类Transformer模型上使用DeepSpeed的激活checkpointing,不能直接封装整个模型,需要针对每个Transformer编码器/解码器层单独应用checkpointing。以下是具体实现步骤:
步骤1:修改LightningModule代码
我们通过包装模型的每个编码器和解码器层,让它们在调用时自动使用DeepSpeed的checkpointing功能:
from lightning.pytorch import LightningModule from transformers import AutoModelForSeq2SeqLM # Flan-T5是seq2seq模型,推荐使用该类 import deepspeed def checkpointed_layer(layer): """包装单个Transformer层,添加DeepSpeed checkpointing支持""" def forward_wrapper(*args, **kwargs): return deepspeed.checkpointing.checkpoint(layer, *args, **kwargs) return forward_wrapper class MyModel(LightningModule): def __init__(self): super().__init__() # 加载Flan-T5-base预训练模型 self.lm = AutoModelForSeq2SeqLM.from_pretrained('google/flan-t5-base') # 对编码器的所有层应用checkpointing for idx in range(len(self.lm.encoder.block)): self.lm.encoder.block[idx] = checkpointed_layer(self.lm.encoder.block[idx]) # 对解码器的所有层应用checkpointing for idx in range(len(self.lm.decoder.block)): self.lm.decoder.block[idx] = checkpointed_layer(self.lm.decoder.block[idx]) def forward(self, x): # 原调用方式保持不变,因为层已被自动包装 return self.lm(**x)
步骤2:配置Trainer的DeepSpeed策略
确保Trainer启用DeepSpeed策略,示例配置如下:
from lightning.pytorch import Trainer trainer = Trainer( accelerator="gpu", devices=1, strategy="deepspeed_stage_2", # 根据硬件情况选择合适的DeepSpeed阶段 max_epochs=3, accumulate_grad_batches=32 )
关键说明
- 为何不封装整个模型?:整个模型的激活数据量过大,直接checkpoint无法有效降低显存占用,反而会因冗余计算拖慢训练速度。针对每层做checkpointing能精准控制激活的存储与重计算。
- 显存节省原理:被checkpoint的层会在前向传播后立即丢弃激活值,反向传播时重新计算这些激活,以少量额外计算开销换取显存空间。
- 适配其他Transformer模型:该方法适用于所有Hugging Face的Transformer架构模型(如BERT、GPT系列),只需调整模型内部的层路径(例如GPT系列的
transformer.h层列表)即可。
内容的提问来源于stack exchange,提问作者BioBroo
相关产品推荐
相关产品推荐

