ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

BigBird省显存利器全解析:梯度检查点recompute_grad与无状态Dropout源码解读

BigBird省显存利器全解析:梯度检查点recompute_grad与无状态Dropout源码解读 BigBird省显存利器全解析梯度检查点recompute_grad与无状态Dropout源码解读【免费下载链接】bigbirdTransformers for Longer Sequences项目地址: https://gitcode.com/gh_mirrors/bi/bigbirdBigBird 是一个基于稀疏注意力sparse-attention的长序列 Transformer 实现官方定位Transformers for Longer Sequences专为 BERT、Pegasus 等模型处理 4096 甚至更长的文档而设计。长序列训练的天敌是显存/HBM 占用而 BigBird 在core目录里内置了两招省显存利器梯度检查点recompute_grad用重算换显存和无状态 Dropout让重算时的随机掩码可复现。本文带你从源码角度看懂它们是如何配合、以极小的速度代价大幅压缩激活内存的 一、为什么 BigBird 必须省显存长序列场景下Transformer 每一层都要保存大量中间激活值供反向传播使用。序列从 512 拉长到 4096激活内存近似呈平方级膨胀GPU/TPU 很容易 OOM。BigBird 用块稀疏注意力把注意力计算量降下来了下方基准图中它明显比同类长上下文模型更省内存且不损精度但激活内存依然可观于是引入经典的梯度检查点思路前向传播时不保留每层的中间激活反向传播时只重算一次前向来拿到激活再求梯度——以约 30%~33% 的额外计算换回大比例的激活显存。 上图正是官方 README 中的基准对比BigBird 在六大长上下文任务上内存消耗显著低于同类模型且精度不打折。二、三步开启梯度检查点最快上手路径整套机制由一个布尔开关驱动三步即可生效打开开关bigbird/core/flags.py中定义了use_gradient_checkpointing默认False含义是是否在反向传播时重算编码器前向以省显存。包装编码器/解码器EncoderStackbigbird/core/encoder.py和DecoderStackbigbird/core/decoder.py检测到该开关为真时用工厂函数add_gradient_recomputation()把每一层包一层子类。包一层重算装饰器包装后的层在call中把整层前向函数f交给recompute_grad.recompute_grad(f)之后正常调用即可。核心就这几行bigbird/core/encoder.py的add_gradient_recomputationdef f(layer_input, attention_mask, band_mask, ...): x super(RecomputeLayer, self).call(...) return x f recompute_grad.recompute_grad(f) return f(layer_input, attention_mask, band_mask, ...)对使用者而言零侵入模型代码一行不改开关一开所有层自动进入检查点模式。三、recompute_grad 源码解读重算是怎么做的核心文件bigbird/core/recompute_grad.py孵化了一个XLA 兼容版的tf.recompute_grad四个部件各管一摊RecomputeContext重算上下文一个线程局部上下文栈_ContextStack每个被包装的层压入一个上下文记录两件事——is_recomputing当前是否处于重算阶段和seed随机数种子。嵌套调用时通过children队列维护父子关系支持任意层嵌套。前向阶段只执行一次f记下上下文但不保留中间激活这就是省显存的关键。反向阶段grad 函数用相同的输入和相同的 seed把f完整重算一遍在新开的GradientTape里求梯度。因为 seed 相同重算出的激活与首次前向完全一致梯度才正确。XLA/TPU 兼容XLA 编译时会忽略控制依赖control dependency执行顺序可能乱序。_in_xla_context()检测到 XLA 环境后_force_data_dependency()会通过一个极小的浮点加法构造假数据依赖强制重算严格发生在梯度流入之后保证顺序正确。 一句话总结上下文管种子、前向只管算、反向再算一遍、XLA 下靠数据依赖锁顺序。四、无状态 Dropout重算的配套件梯度检查点有个隐藏坑如果层里用了普通 Dropout反向时重算前向会重新掷一次骰子掩码和第一次不一样重算出的激活就对不上了梯度直接错掉。BigBird 的解法是把 Dropout 变成纯函数recompute_grad.py后半部分stateless_dropout()不依赖任何全局随机状态必须显式传入seed内部用tf.random.stateless_uniform采样。同样的输入 同样的 seed ⇒ 同样的掩码天然可重放。RecomputingDropout继承 KerasLayer的智能切换层。每次前向时先查get_recompute_context()处于重算上下文 → 用tf.stack([上下文seed, 本层专属seed])调用stateless_dropout保证可复现不在上下文里比如推理→ 退回普通tf.nn.dropout行为不变。每层初始化时都会生成一个随机的_recompute_seed与全局上下文 seed 组合保证层与层掩码不同、同层重算一致两全其美。五、在项目里实际用在哪注意力层bigbird/core/attention.py中attention_dropout即RecomputingDropout编码器bigbird/core/encoder.py中每个层都有attention_dropout/output_dropout两个实例解码器bigbird/core/decoder.py中自注意力、交叉注意力、输出各有一个RecomputingDropout官方 Notebookbigbird/classifier/imdb.ipynb、bigbird/summarization/pubmed.ipynb在训练前直接设置FLAGS.use_gradient_checkpointing True即可开启训练脚本如bigbird/classifier/base_size.sh顶部注释推荐TF_XLA_FLAGS--tf_xla_auto_jit2这正是_in_xla_context()检测的依据之一建议保持开启。六、速查总结组件所在文件作用use_gradient_checkpointingbigbird/core/flags.py总开关一行开启全部检查点add_gradient_recomputationbigbird/core/encoder.py/decoder.py把每层包装为重算层recompute_gradbigbird/core/recompute_grad.py核心不存激活、反向重算、XLA 兼容stateless_dropoutbigbird/core/recompute_grad.py种子驱动的确定性 DropoutRecomputingDropoutbigbird/core/recompute_grad.py上下文感知的自动切换 Dropout 层核心思想显存不够重算来凑——只要所有随机操作Dropout都改成给定种子即可复现的无状态形式重算就是完全等价的。这套组合拳对任何长文档摘要、长文本预训练场景哪怕换到自己的项目里都通用是长序列大模型训练的必备技能之一。【免费下载链接】bigbirdTransformers for Longer Sequences项目地址: https://gitcode.com/gh_mirrors/bi/bigbird创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表