已关闭
[RFC] HyperParallel 支持重计算和swap统一管理 #5
DavidFFFan创建于 1月19日关闭于 6月23日
1月19日 创建了Requirement
1月19日 关联了MindSpore/hyper-parallel Pull Request !115
1月19日 修改了标题
1月19日 修改了描述
1月19日 修改了描述
此处折叠了90条消息 查看更多
5月12日 修改了issue 的描述
5月12日 修改了issue 的描述
6月23日 issue状态由 TODO 改变为 DONE
6月23日 关闭了 issue
MindSpore-Bot
6月23日 评论:
6月23日 评论:


背景与目标描述.
在深度神经网络训练过程中(尤其是 Transformer 等大模型),反向传播需要依赖前向阶段产生的中间激活值(Activation)来计算梯度。
随着模型规模和层数的不断增加(例如上百层 Transformer),这些激活值会占用大量 device 显存(GPU / NPU),成为训练过程中的主要瓶颈。
重计算用途:
重计算(Activation Checkpointing) 通过“用时间换空间”的方式,在前向传播阶段不保存中间激活值,而是在反向传播阶段按需重新执行前向计算以恢复激活值,从而显著降低显存占用。
约束:
1.在实际模型中仍存在一类问题算子:
2.Checkpoint 输入本身也是显存占用的重要来源,重计算对降低显存效果不佳。
3.对于swap功能,有些跨层共享内存的tensor不能被swap。
建议的方案.
方案概述:
引入 Swap 机制 作为重计算策略的补充能力,对 Checkpoint 相关激活数据 进行统一管理,包括:
1.算子输出(Operator Output)
2.Checkpoint 输入(Checkpoint Input)
核心思路为:
在前向阶段:
在重计算阶段:
通过 异步 IO:
声明式接口:
定义于hyper-parallel\hyper_parallel\core\activation_checkpoint\activation_checkpoint.py的声明式接口用法如下:
选择重计算 - checkpoint_wrapper
只开重计算
重计算+SWAP
只SWAP激活 - swap_wrapper
自定义policy_fn:
checkpoint_wrapper的policy_fn函数签名:
ctxSelectiveCheckpointContext对象,self.is_recompute识别前向反向op*args, **kwargs返回值(
CheckpointPolicy枚举):MUST_SAVE/PREFER_SAVEMUST_RECOMPUTE/PREFER_RECOMPUTEMUST_SWAPSwapManager使用)swap_wrapper的policy_fn函数签名:
tensor返回值(
CheckpointPolicy枚举):CheckpointPolicy.MUST_SAVECheckpointPolicy.MUST_SWAPSwapManager使用)基于DispatchMode的选择重计算:
定义于hyper-parallel\hyper_parallel\platform\mindspore\activation_checkpoint\sac.py的 create_selective_checkpoint_contexts,本质上是在给重计算提供一对“前向态 / 反向态”的上下文拦截。
它返回一对 context:
_CachingMindSporeDispatchMode(...)
_CachedMindSporeDispatchMode(...)
完整调用链
forward_ctx, recompute_ctx = context_fn()
前向传播 — forward_ctx
forward_ctx继承 MsDispatchMode(MindSpore 的算子分发拦截机制),其 ms_dispatch 对 block 内每个算子都会触发。执行算子后根据policy_fn策略决定是否缓存输出(保存到内存-SAVE,卸载到CPU-SWAP,不缓存-RECOMPUTE)。
同时,外层的_CreatePlaceholderHook (MS提供能力)拦截 block 内所有 saved_tensors_hooks(即 autograd 引擎自动 save 的每个 tensor),不真正存 tensor,改为存一个_PlaceHolder 对象,block 内所有中间激活都被"清空",反向时一律触发重计算。两层机制的协同:
• 外层粗粒度(block 级):所有的 saved tensor 替换 → 触发 recompute
• 内层细粒度(算子级):forward_ctx 选择性缓存昂贵算子的输出
反向传播 — recompute_ctx
当_PlaceHolder 被 unpack 时,触发 recompute_function:
recompute_ctx 在此拦截所有算子分发:
if policy in (MUST_SAVE, PREFER_SAVE) # 直接从 storage 取回前向缓存的输出,不执行算子!
if policy == MUST_SWAP # 取回已从 CPU 加载回来的 tensor
if policy in (MUST_RECOMPUTE / PREFER_RECOMPUTE) # 真正重计算:执行算子
选择重计算的卸载inputs能力:
利用了
saved_tensors_hooks的不可嵌套性,将 inputs 和中间激活的处理精确地分离开来:_InputSaver.apply()调用期间AsyncSaveOnCpuctx.save_for_backward(*inputs)→ 由 hook 将 inputs 卸载到 CPU_CreatePlaceholderHook开启后_CreatePlaceholderHookSwap 覆盖范围:
Swap 机制统一管理 Checkpoint Activation,覆盖以下两类对象:
触发时机区别如下:
Swap 管理:
系统采用四层抽象:
classDiagram direction TB class SwapTensor { -val -val_cpu -_state -is_slice_tensor -storage_size +get_val() +async_load() +wait_load() +async_offload() +wait_offload() +state } class Storage { -save_storage : Dict[Any, List[Any]] -swap_storage : Dict[Any, List[Any]] +launch_load() +wait_load() +launch_offload() +wait_offload() } class SwapGroup { -group_name : str -_storages : WeakSet[Storage] -_load_event -_offload_event +add(storage) +launch_offload(copy_stream) +wait_offload() +launch_load(copy_stream) +wait_load() } class SwapManager { <<Singleton>> -_groups : Dict[str, SwapGroup] -_current_group_name : str -_copy_stream -_layer_count : int +add_storage(group_name, storage) +launch_offload(group_name, copy_stream) +wait_offload(group_name) +launch_load(group_name, copy_stream) +wait_load(group_name) +set_forward_prefetch_layer(first_layer, second_layer) } SwapManager "1" o-- "*" SwapGroup : manages SwapGroup "1" o-- "*" Storage : contains Storage "1" o-- "*" SwapTensor : wraps异步 IO 设计:
非PP场景(按层)
针对常见模型结构,提供按层的高阶接口
SwapManager().set_forward_prefetch_layer(first, second)配置layer正向的执行顺序,从而自动管理 offload / load 时机,实现计算与拷贝并发。from hyper_parallel.core.activation_checkpoint import CheckpointPolicy, SwapManager, checkpoint_wrapper op_non_recompute = { torch.ops.aten.matmul.default, torch.ops.aten.addmm.default, torch.ops.aten.bmm.default } def policy_fn(ctx, op, *args, **kwargs): if op in op_non_recompute: return CheckpointPolicy.MUST_SWAP return CheckpointPolicy.MUST_RECOMPUTE for i, layer in enumerate(model.layers): model.layers[i] = checkpoint_wrapper(layer, policy_fn=policy_fn, swap_inputs=True) for i in range(len(model.layers) - 1): SwapManager().set_forward_prefetch_layer(model.layers[i], model.layers[i + 1])复杂场景(Pipeline Parallel 等)
为支持复杂并行拓扑,提供 原子级 Swap 接口,由用户或调度器精确控制 swap 时序,可以实现按照stage/chunk等粒度调度。
# 设置 swap 分组名称,用于限定 swap 范围 SwapManager().set_current_group_name(group_name) # 异步 offload / load 接口,支持自定义 stream SwapManager().launch_offload(group_name, copy_stream=None) SwapManager().wait_offload(group_name) SwapManager().launch_load(group_name, copy_stream=None) SwapManager().wait_load(group_name)通过分组机制实现:
涉及到的对外API
checkpoint_wrapper
功能:
给一个 module/func 加上 activation checkpoint 能力。
当传了 policy_fn,则进入 selective checkpoint 模式,可对不同算子决定:保存,重算,卸载。
当传 swap_inputs=True,会把 checkpoint 区域里 autograd 保存的输入 tensor 走 async_save_on_cpu 逻辑,卸 载到host进一步降显存。
入参:
module/func:被包装对象。
policy_fn=None:用来决定 selective checkpoint 策略。
swap_inputs=False:是否把 checkpoint 保存的输入也做 CPU 异步换出。
出参:返回一个“包装后的模块对象”。
swap_wrapper
功能:给一个 module/func 加上 activation swap 能力。
入参:
module/func:被包装对象。
policy_fn=None:可选策略函数。
出参:返回一个“包装后的模块对象”。
set_forward_prefetch_layer
功能:把两个相邻层注册成前后关系,为这两个层建立 swap group,并注册 forward/backward hooks,作用是让 swap 形成时序流水。
入参:
self:SwapManager() 单例对象。
first_layer:前一个swap组,通常是 model.layers[i]。
second_layer:后一个swap组,通常是 model.layers[i + 1]。
出参:无。
set_current_group_name
功能:设置当前正在执行 forward 的 SwapGroup 名称,用于在单例 SwapManager 上记录当前活跃的 swap 分组。
入参:
self:SwapManager() 单例对象。
group_name:swap group名。
出参:无。
launch_offload
功能:异步启动指定 swap group 中所有张量的 D2H 数据搬运。
入参:
self:SwapManager() 单例对象。
group_name:目标 swap group 的名称。
copy_stream (可选:用于执行数据搬运的流;缺省时使用单例_copy_stream。
出参:无。
wait_offload
功能:阻塞等待指定 swap group 的 offload 完成。
入参:
self:SwapManager() 单例对象。
group_name:目标 swap group 的名称。
出参:无。
launch_load
功能:异步启动指定 swap group 中所有张量的 H2D 数据预取。
入参:
self:SwapManager() 单例对象。
group_name:目标 swap group 的名称。
copy_stream (可选:用于执行数据搬运的流;缺省时使用单例_copy_stream。
出参:无。
wait_load
功能:阻塞等待指定 swap group 的 load 完成。
入参:
self:SwapManager() 单例对象。
group_name:目标 swap group 的名称。
出参:无。
开发自验
1.ckpt_activation.py::test_ac_memory_comparison
对比 7 种模式的训练结果和显存峰值:none、recompute、funcrecompute、save、funcsave、swap、funcswap。
每种模式都在独立子进程里跑 3 个 step。
校验所有模式在每个训练 step 的 loss 和基线 none 一致。
校验显存趋势:mem_none > mem_recompute ≈ mem_swap。
2.swap_activation.py::test_act_swap_memory_comparison
对比 3 种模式:none、swap、swap_with_policy。
每种模式在独立子进程里训练 3 个 step。
校验每个 step 的 loss 与基线 none 一致。
校验显存层级关系:mem_none > mem_swap_with_policy > mem_swap
测试范围
重计算/SWAP动态图在非并行场景下验证
重计算/SWAP动态图结合并行场景验证
细粒度选择重计算验证
1.需验证swap_wrapper和checkpoint_wrapper的selective policy 混合策略,整层包装与子模块(func)包装。必须满足训练无异常,与同模型、同输入、同随机种子、同优化器、同 step 数的baseline的loss对齐。
2.checkpoint_wrapper + MUST_RECOMPUTE模式,checkpoint_wrapper + MUST_SWAP模式,swap_wrapper模式峰值显存应下降。