ARTICLE DETAIL

资讯详情

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

TorchTitan 的 torch_remat 区域激活检查点:RegionAC 的配置与模型集成实践

TorchTitan 的 torch_remat 区域激活检查点:RegionAC 的配置与模型集成实践 TorchTitan 的 torch_remat 区域激活检查点:RegionAC 的配置与模型集成实践【免费下载链接】torchtitanA PyTorch native platform for training generative AI models项目地址: https://gitcode.com/GitHub_Trending/to/torchtitan本文围绕 TorchTitan 仓库中的 docs/remat.md 展开,系统讲解基于torch_remat的区域激活检查点(Region AC):它是如何作为全量/选择性激活检查点的模型感知替代方案,如何在模型代码中声明保存区域(save regions)与重计算依赖,以及训练配置应如何约束随机状态。读完后,你将能够读懂RegionAC的完整调用链,为自己实现的 transformer 模块正确添加remat.region标注,并理解仓库现有模型(Attention、FFN、MoE、视觉编码器)中每一处区域标注的设计依据。一、动机:从算子级检查点到模型感知的区域检查点TorchTitan 目前提供多种激活检查点(Activation Checkpointing, AC)策略,统一实现在 activation_checkpoint.py 中,均继承自ActivationCheckpointing基类(见torchtitan/distributed/activation_checkpoint.pyL118-L170):FullAC(L172-L188):用ptd_checkpoint_wrapper包裹整个 transformer block,反向传播时整块重计算,省内存但计算开销最大;SelectiveAC(L191-L297):在算子级做保存决策,基于create_selective_checkpoint_contexts的自定义 policy,保存重计算昂贵的算子(各类 matmul、SDPA、通信算子等,见_get_default_save_ops(),L36-L100),其余重计算;RegionAC(L302-L373):torch_remat提供的模型感知替代方案。docs/remat.md对这一动机的表述是:选择性 AC 的保存决策发生在算子(aten op)层面,而torch_remat让模型代码标注语义区域,训练配置决定保留哪些区域的输出。这带来两个直接好处:策略可见性:保存/重计算策略以显式注解的形式出现在模型源码里,紧贴它控制的操作,而不是一份与代码分离的算子白名单;配置指向稳定的模型概念:配置引用的是 attention projection(如attention.qkv)这类稳定的模型概念,而不是某个具体的 ATen 算子符号。源码中RegionAC的 docstring 也印证了文档的定位——它要求模型声明兼容的区域并提供显式保存策略,同时注释中保留了迁移路线(见activation_checkpoint.pyL300-L301):# TODO: Migrate the existing AC implementations to RegionAC and keep RegionAC # as the single activation-checkpointing implementation.这与文档长期计划是把现有FullAC和SelectiveAC迁移到torch_remat,收敛为唯一的 AC 实现一致;文档同时提醒RegionAC是临时命名,可能随迁移进展改变。二、配置保存区域:RegionAC.Config.save_regions2.1 基本用法保存区域由save_regions字段配置,内容是相对于单个 transformer block 的 shell 风格 glob 模式(docs/remat.md给出的示例):RegionAC.Config( save_regions[ attention.qkv, attention.wo, ] )语义规则如下:命中模式(含attention.*这类通配符)的区域,其输出在反向传播时保留;未命中的区域全部重计算——即默认倾向是重算,保存是显式选择;同一策略适用于所有 transformer block:区域名相对于 block 而非模型,因此attention.qkv会在每一层的 attention 子模块上生效。源码注释明确暂不支持按 block 区分策略(见RegionAC.Config.save_regions字段文档,activation_checkpoint.pyL312-L322);未匹配任何区域的模式目前被静默忽略。文档指出验证逻辑必须最终考虑跨 pipeline stage 的区域分布,源码中对应一条 TODO(见activation_checkpoint.pyL363-L364),说明在多阶段流水线并行下,某个模式可能只在别的 stage 的 block 中才存在,当前校验无法覆盖这一情况。2.2 配置层约束:两个必须知道的限制RegionAC.Config覆写了__post_init__(activation_checkpoint.pyL330-L340),在构造时强制两条约束,并有单测锁定这些报错信息(test_remat.py 的test_unsupported_config_options_error):preserve_rng_state必须为False。RegionAC.Config将基类中默认True的preserve_rng_state改为默认False,一旦显式置True会抛出ValueError,提示改用torch_remat.RecomputeStateHook管理随机状态;不支持debug选项:激活检查点的 debug 抓取功能与torch_remat不兼容,置True同样抛错。此外,save_regions是必填字段——没有默认值也没有default_factory,单测test_save_regions_config_is_required直接断言该字段为dataclasses.MISSING,确保选择 RegionAC 时用户必须做出显式的保存决策。2.3 策略选择入口RegionAC与其他 AC 策略通过 tyro 子命令暴露为 Trainer 配置的联合类型(activation_checkpoint.pyL419-L428):ActivationCheckpointingConfig ( Annotated[SelectiveAC.Config, tyro.conf.subcommand(selective)] | Annotated[RegionAC.Config, tyro.conf.subcommand(region)] | Annotated[FullAC.Config, tyro.conf.subcommand(full)] | Annotated[MemoryBudgetAC.Config, tyro.conf.subcommand(memory-budget)] | Annotated[None, tyro.conf.subcommand(none)] )即五个可选子命令:selective、region、full、memory-budget、none。选择region即启用本文主题。各模型(如 llama3、qwen3 的parallelize函数)以及实验性 graph trainer 的并行化入口都以ac_config: ActivationCheckpointingConfig参数接收该配置,统一在并行化阶段对模型施加拉展。三、在模型代码中声明区域3.1 注解语法模型代码在被控制的操作处定义区域(docs/remat.md示例):q, k, v remat.region( self.qkv_linear, self.remat_region_name(qkv), recomputeself.remat_should_recompute(qkv), )(x)三个要素对应torchtitan/protocols/module.py中Module基类提供的辅助方法(L53-L84):remat_region_name(local_name)(L56-L60):把局部名(如qkv)解析为限定名(如attention.qkv)。它读取_remat_module_fqn(该模块相对于 transformer block 的完整模块名)。未配置时退化为返回局部名,保证模型在未启用 RegionAC 时行为不变;remat_should_recompute(local_name)(L62-L67):用fnmatch把限定名与_remat_save_patterns逐一匹配,只要命中任一保存模式就返回False(保留输出),否则返回True(重计算);configure_remat_regions(save_patterns)(L69-L84):由RegionAC.apply()在模型树递归调用,把每个Module子模块的_remat_module_fqn设为其相对 block 的 fqn,并写入统一的保存模式元组。一个关键设计:在没有外层remat.checkpoint时,remat.region不改变任何执行行为。Module类注释也点明了这一点(module.pyL51-L52):这些默认值被 RegionAC 在每个 checkpointed block 的所有 Module 上替换,但在包围性的 torch_remat checkpoint 之外,它们不影响执行。这意味着模型代码中的区域注解是惰性的——只有选择了region子命令、RegionAC实际包裹了 block,注解才开始生效。3.2 仓库中的真实区域标注common 模型组件已在各算子处完成了区域标注,可作为集成范本:前馈层(feed_forward.py L83-L97)——注意w13/w2两个区域之间的recompute_needs_tensor调用,第四节会详细解释:def forward(self, x: torch.Tensor) - torch.Tensor: gate_up_TF remat.region( self.w13, self.remat_region_name(w13), recomputeself.remat_should_recompute(w13), )(x) gate_TF, up_TF gate_up_TF.unflatten(-1, (-1, 2)).unbind(-1) remat.recompute_needs_tensor(gate_TF, up_TF) out_TD remat.region( self.w2, self.remat_region_name(w2), recomputeself.remat_should_recompute(w2), )(self.activation_fn(gate_TF, up_TF)) remat.recompute_needs_tensor(out_TD) return out_TD注意力层(attention.py L911-L950)声明了qkv、inner_attention、wo三个区域,并在 RoPE/attention 掩码处理处按需插入remat.recompute_needs_tensor标记;MoE 路由器(moe.py L288-L296)把 top-k 专家选择包裹进remat.region——从单测test_router_decision_is_always_saved(test_remat.py L395-L418)可以确认:路由决策在检查点内只执行一次(select_experts.call_count 1),保证前向与重计算之间专家分配一致;视觉编码器(vision_encoder.py)标注了w1、w2、qkv、inner_attention、wo五个区域,命名相对 block 前缀为attn.与mlp.;辅助损失(aux_loss.py L191-L196)对由裸操作消费的输出同样补了recompute_needs_tensor标记。各变体 FFN(普通FeedForward、SigmoidGatedFeedForward、DistGEMMFeedForward、fused SwiGLU 变体)声明的区域边界不同(后者多一个gate区域),单测test_feed_forward_variants_use_expected_region_boundaries用remat.collect_trace()追踪实际生成的区域名列表来逐一校验,说明区域命名是可测试、可审计的。四、声明重计算依赖:remat.recompute_needs_tensor这是文档中规则最密集的一节,理解它需要区分两类消费者:消费者是显式的remat.region(..., recomputeTrue):依赖自动推断,无需标记;消费者是普通(裸)算子:torch_remat无法推断,必须在输出被消费处显式声明。规则汇总(均出自docs/remat.mdDeclaring recomputation dependencies一节,并对照仓库实现验证):放置位置在消费者侧:标记必须放在即将读取该张量的裸操作之前,而不是产生该张量的区域之后。这样输出只在消费者真正运行时才被保留;多个区域输出被同一个裸算子消费时,全部传给同一个调用,不同消费者保持各自独立的调用。仓库范本即 FFN 中的remat.recompute_needs_tensor(gate_TF, up_TF),以及SigmoidGatedFeedForward里对两个区域输出out_TD、gate_out_TD的一次性标记(feed_forward.py L115-L123);消费者是另一个remat.region时不要加标记,依赖会被自动推断;可以省略的情况:区域输出直接从被检查点的 transformer block 返回,且 block 内没有任何操作读取其数据。位于子模块中的最后一个区域并不足够——例如 attention 输出仍会被外层 transformer block 的残差加法消费,此时必须标记;如果消费者位于拥有该区域的 helper 模块之外,应把标记放到模块边界允许下尽可能贴近该调用点的位置;可以传视图(view):torch_remat按 storage 把视图解析回产生它的区域,因此unbind/unflatten得到的子张量可直接传入,如上例的gate_TF、up_TF。为什么这一步如此关键?文档的解释是:在原始前向中,torch_remat负责判断保存区域的输出在重计算时是否需要;若缺少标记,一个被普通重计算操作所需的张量可能不会被保留,反向传播时就会出错或产生错误梯度。五、随机状态管理RegionAC要求preserve_rng_stateFalse。含义是:不允许用传统 checkpoint 那种stash restore RNG 状态的机制来保证重计算时的随机一致性,任何可能在保存区域内推进的随机状态,必须改用显式的torch_remat.RecomputeStateHook管理。这一约束在三个层面被强制:配置层:RegionAC.Config的__post_init__对preserve_rng_stateTrue直接抛错(见 2.2 节);实现层:_wrap_block中remat.checkpoint的调用硬编码preserve_rng_stateFalse(activation_checkpoint.pyL347-L351),即使配置试图放宽也无法生效;文档层:docs/remat.md末尾专门用Random state一节重申该要求。对集成者的实际影响:如果你的保存区域内包含 dropout 等消费随机数的操作,不能依赖自动状态保存,而需要为这些操作注册RecomputeStateHook来显式控制前向与重计算之间的随机状态衔接;若区域内不含随机操作,则无需任何额外处理。六、运行机制:RegionAC如何作用于模型从源码看RegionAC的完整应用流程(activation_checkpoint.pyL355-L373 L342-L353):apply(model)先取model.get_submodule(layers)下的所有直接子模块作为 transformer block;若无 block,记录日志并直接返回(适配流水线并行下可能不含层的模型分片);对每个 block,先调用configure_remat_regions(save_regions)(见第三节),把区域限定名解析规则与保存模式写入模块树;再调用_wrap_block:以 block 的 fqn(如layers.0)作为region_name,把module.forward包装成remat.checkpoint(region_name..., determinism_check..., preserve_rng_stateFalse)的结果,即每个 block 的前向整体成为一个torch_remat检查点区域,内部再由remat.region划分出可保留输出的子区域。与FullAC/SelectiveAC的外层ptd_checkpoint_wrapper包裹整个 block不同,RegionAC不引入 wrapper 模块,而是直接替换forward引用——单测test_llama_attention_policy_applies_without_changing_state_dict(test_remat.py L227-L236)专门验证了应用RegionAC后state_dict的键完全不变,即该策略对 checkpoint 格式透明,不影响已有的权重加载与转换流程。七、测试给出的行为契约test_remat.py 用计数操作(_CountingOp记录前向次数) 数值严格比对的方式,把文档描述的语义固化成可执行断言,值得集成时对照:保存策略save_regionsqkv/inner_attention/wo的前向次数[](全部重算)(2, 2, 2)[attention.*](1, 1, 1)[attention.qkv](1, 2, 2)[attention.inner_attention](2, 1, 2)[attention.wo](2, 2, 1)每次运行的前向输出、输入梯度与参数梯度都与未启用检查点的基线逐位相等(rtol0, atol0),验证了保留区域不改变数值,只改变重计算范围这一核心契约。FFN(w13/w2)与视觉 block(attn.qkv、mlp.w1等)有同构的用例,包括[attn.*, mlp.*]多通配符组合。八、当前限制与路线区域模式目前跨 block 统一,不支持按层差异化策略(源码字段注释明确Per-block remat policies are not currently supported);未匹配模式暂被忽略,跨 pipeline stage 的模式校验是待完成项(源码 TODO 文档说明);不支持debug选项与preserve_rng_stateTrue(构造期即报错);RegionAC命名为临时,文档与源码 TODO 均指向同一目标:把FullAC、SelectiveAC迁移到torch_remat,收敛为单一激活检查点实现。参考路径汇总内容路径官方设计文档docs/remat.mdAC 策略统一实现与RegionACtorchtitan/distributed/activation_checkpoint.pyModule的区域辅助方法与配置注入torchtitan/protocols/module.pyFFN 区域标注范本torchtitan/models/common/feed_forward.pyAttention 区域标注torchtitan/models/common/attention.py行为契约单测tests/unit_tests/gpu/test_remat.py【免费下载链接】torchtitanA PyTorch native platform for training generative AI models项目地址: https://gitcode.com/GitHub_Trending/to/torchtitan创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表