FiD显存优化秘籍:Checkpointing与answer_maxlength如何驯服100段长文本
FiD显存优化秘籍Checkpointing与answer_maxlength如何驯服100段长文本【免费下载链接】FiDFusion-in-Decoder项目地址: https://gitcode.com/gh_mirrors/fi/FiDFiDFusion-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_lengthTrue并开启truncation默认值为-1表示不截断src/options.py这正是变长的默认状态。所以训练 100 段长文本时把它设为一个合理值例如 50与推理时 generate 的 max_length50 对齐就能把解码器张量钉死显著降低显存峰值。四、一步到位官方 large 读者的完整参数官方用 64 张 GPU 训练 t5-large 版 FiD 时就是同时启用这两个开关并配合per_gpu_batch_size 1README.mdpython 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 basebase 比 large 省数倍显存入门首选缩短段落--text_maxlength直接线性降低编码器与交叉注意力的显存梯度累积--accumulation_steps小 batch 累积步数等效放大 batchsrc/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 EMREADME.md依赖注意项目基于 PyTorch 1.6 与 Transformers3.0.2README.md版本不匹配容易踩坑总结两行参数显存减半✅--use_checkpoint梯度检查点重算激活砍掉编码器侧最大头的显存 ✅--answer_maxlength给解码器定长消灭变长张量的碎片与峰值。再加上per_gpu_batch_size 1 多卡并行这套组合拳普通集群也能稳稳训练 100 段长文本的 FiD 大模型。显存不够先从这两个开关查起吧。【免费下载链接】FiDFusion-in-Decoder项目地址: https://gitcode.com/gh_mirrors/fi/FiD创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考