FiD显存优化秘籍:Checkpointing与answer_maxlength如何驯服100段长文本
FiD显存优化秘籍:Checkpointing与answer_maxlength如何驯服100段长文本
【免费下载链接】FiDFusion-in-Decoder项目地址: https://gitcode.com/gh_mirrors/fi/FiD
FiD(Fusion-in-Decoder,解码器融合)是开放域问答领域的经典生成式模型,一次要"读完"100个检索段落再作答,显存压力巨大。本文带你掌握 FiD 显存优化的两大核心手段——梯度检查点(--use_checkpoint)与答案长度固定(--answer_maxlength),教你用有限显卡驯服 100 段长文本的训练任务。
一、为什么 FiD 训练 100 段长文本会"吃掉"显存?
FiD 的巧妙之处在于:它用一个 T5 编码器并行处理 100 个段落(每个问题 + 段落拼接成一条输入),再让解码器通过交叉注意力在全部 100 段拼接后的长序列上"融合"信息生成答案。模型定义见 FiDT5。
这意味着显存占用随段落数线性增长:
- 编码器侧:输入被 reshape 成
(batch × 100) × 250的张量,激活值(activations)规模同样放大 100 倍; - 解码器侧:交叉注意力的 Key/Value 长度是
100 × text_maxlength,注意力矩阵也随之膨胀。
论文作者的官方说明也很直白:「用 100 个段落训练这些模型非常吃显存,我们通过 checkpointing 来缓解这一问题」(原文见 README.md)。下面两个"开关"正是为此而生。
二、秘籍①:--use_checkpoint,用时间换空间
🔥梯度检查点(Gradient Checkpointing)的思想很简单:前向传播时不保存每个编码层的中间激活,反向传播时重新计算一遍。代价是多约 1/3 的前向计算时间,收益是激活显存从"保存所有层"骤降到"只保存检查点层"。
FiD 的实现集中在 src/model.py 中,思路分三步:
- 包装编码器:
wrap_encoder()用 EncoderWrapper 把 T5 编码器包起来,训练时把 100 段"压平"成一个大 batch 处理,结束后再恢复形状; - 逐层加装检查点:apply_checkpoint_wrapper 把编码器的每一层都包进 CheckpointWrapper,其中真正调用
torch.utils.checkpoint.checkpoint的地方就是它; - 动态开关:set_checkpoint() 在训练入口(train_reader.py)根据命令行参数一键启停。
一个贴心的细节:CheckpointWrapper只在self.training为真时才启用重计算,所以推理阶段(test_reader.py 生成答案时)完全不受拖累,速度不受影响。
三、秘籍②:--answer_maxlength,给解码器"定长"
如果说 checkpointing 优化的是编码器,那么--answer_maxlength针对的是解码器侧的"变长张量"问题。
📏 编码器输入的长度是固定的(text_maxlength控制),但解码器要学习的目标答案长短不一:有的答案是 3 个 token,有的接近 50 个。变长张量会导致:
- 每个 batch 分配大小不一的显存块,产生内存碎片和峰值开销;
- 分布式多卡训练时,各卡形状不一致还会带来同步麻烦。
解决方法就是在数据整理阶段把答案统一补齐/截断到固定长度。这一步发生在 Collator 里:
answer_maxlength > 0时,token 化会pad_to_max_length=True并开启truncation;- 默认值为
-1,表示不截断(src/options.py),这正是"变长"的默认状态。
所以训练 100 段长文本时,把它设为一个合理值(例如 50,与推理时 generate 的 max_length=50 对齐),就能把解码器张量"钉死",显著降低显存峰值。
四、一步到位:官方 large 读者的完整参数
官方用 64 张 GPU 训练 t5-large 版 FiD 时,就是同时启用这两个开关,并配合per_gpu_batch_size 1(README.md):
python train_reader.py \ --use_checkpoint \ --answer_maxlength 50 \ --lr 0.00005 \ --optim adamw \ --scheduler linear \ --weight_decay 0.01 \ --text_maxlength 250 \ --per_gpu_batch_size 1 \ --n_context 100 \ --total_step 15000 \ --warmup_step 1000💡 参数含义速查(定义见 src/options.py):
--n_context 100:每个问题配 100 个上下文段落;--text_maxlength 250:每段(问题+段落)最多 250 个 token;--per_gpu_batch_size 1:单卡 batch 为 1,靠多卡堆吞吐;--use_checkpoint:启用第二节的梯度检查点。
五、更多省显存技巧清单
| 技巧 | 参数 | 说明 |
|---|---|---|
| 换小模型 | --model_size base | base 比 large 省数倍显存,入门首选 |
| 缩短段落 | --text_maxlength | 直接线性降低编码器与交叉注意力的显存 |
| 梯度累积 | --accumulation_steps | 小 batch + 累积步数,等效放大 batch(src/options.py) |
| 多卡/多机 | local_rank+ SLURM | 分布式拆分数据,多机流程见 src/slurm.py |
| 推理省显存 | 无需额外配置 | 检查点只在训练时生效,推理天然轻量 |
六、快速上手:5 分钟跑通 FiD
# 1. 获取代码 git clone https://gitcode.com/gh_mirrors/fi/FiD cd FiD # 2. 下载数据与预训练模型(脚本见仓库根目录) bash get-data.sh bash get-model.sh -m nq_reader_base- 训练入口:train_reader.py
- 评测入口:test_reader.py,官方 base 模型在 NaturalQuestions 上可达 50.1 EM(README.md)
- 依赖注意:项目基于 PyTorch 1.6 与 Transformers3.0.2(README.md),版本不匹配容易踩坑
总结:两行参数,显存减半
✅--use_checkpoint:梯度检查点重算激活,砍掉编码器侧最大头的显存; ✅--answer_maxlength:给解码器定长,消灭变长张量的碎片与峰值。
再加上per_gpu_batch_size 1+ 多卡并行这套组合拳,普通集群也能稳稳训练 100 段长文本的 FiD 大模型。显存不够?先从这两个"开关"查起吧。
【免费下载链接】FiDFusion-in-Decoder项目地址: https://gitcode.com/gh_mirrors/fi/FiD
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
