已关闭
[RFC] HyperParallel 支持重计算和swap统一管理 #5
DavidFFFan创建于  1月19日关闭于  6月23日
DavidFFFan
DavidFFFan成员
1月19日 创建

背景与目标描述.

在深度神经网络训练过程中(尤其是 Transformer 等大模型),反向传播需要依赖前向阶段产生的中间激活值(Activation)来计算梯度。
随着模型规模和层数的不断增加(例如上百层 Transformer),这些激活值会占用大量 device 显存(GPU / NPU),成为训练过程中的主要瓶颈。

重计算用途:

重计算(Activation Checkpointing) 通过“用时间换空间”的方式,在前向传播阶段不保存中间激活值,而是在反向传播阶段按需重新执行前向计算以恢复激活值,从而显著降低显存占用。

约束:

1.在实际模型中仍存在一类问题算子:

  • 计算开销较大(如大规模 MatMul、Attention)
  • 重计算代价高,显著影响整体性能
  • 并不适合采用重计算策略

2.Checkpoint 输入本身也是显存占用的重要来源,重计算对降低显存效果不佳。

3.对于swap功能,有些跨层共享内存的tensor不能被swap。


建议的方案.

方案概述:

引入 Swap 机制 作为重计算策略的补充能力,对 Checkpoint 相关激活数据 进行统一管理,包括:
1.算子输出(Operator Output)
2.Checkpoint 输入(Checkpoint Input)

Torch原生 Hyper-Parallel
torch原生能力 hyper-parallel增强

核心思路为:

  • 在前向阶段:

    • 对选定的 算子的输出 和 CheckPoint输入 加入到SwapManager中统一管理
    • 正向完成后,触发SwapManager的异步offload
  • 在重计算阶段:

    • 反向执行前通过SwapManager预取下一层反向数据;
    • 反向执行前等待当前层数据预取完成。
  • 通过 异步 IO:

    • 实现计算与拷贝并发
    • 降低 swap 对整体性能的影响

声明式接口:

定义于hyper-parallel\hyper_parallel\core\activation_checkpoint\activation_checkpoint.py的声明式接口用法如下:

选择重计算 - checkpoint_wrapper

只开重计算

from hyper_parallel.core.activation_checkpoint import checkpoint_wrapper

for i, layer in enumerate(model.layers):
    model.layers[i] = checkpoint_wrapper(layer)

重计算+SWAP

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
}

# 策略函数,用于配置哪些正向算子做SWAP/SAVE/RECOMPUTE
def policy_fn(ctx, op, *args, **kwargs):
    if op in op_non_recompute:
        return CheckpointPolicy.MUST_SWAP # 算子输出配置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) # swap_inputs用于控制checkpoint输入的swap。

# 高阶接口,用于配置offload和prefetch;
# 如果不调用,不会触发swap,行为和torch SAVE策略相同。
for i in range(len(model.layers) - 1):
    SwapManager().set_forward_prefetch_layer(model.layers[i], model.layers[i + 1])
只SWAP激活 - swap_wrapper
from hyper_parallel.core.activation_checkpoint import SwapManager
for i, layer in enumerate(model.layers):
    model.layers[i].attn = swap_wrapper(layer.attn)

for i in range(len(model.layers) - 1):
    SwapManager().set_forward_prefetch_layer(model.layers[i], model.layers[i + 1])

自定义policy_fn:

checkpoint_wrapper 的 policy_fn

函数签名:

policy_fn(ctx: SelectiveCheckpointContext, op, *args, **kwargs) -> CheckpointPolicy
参数 说明
ctx SelectiveCheckpointContext 对象,self.is_recompute识别前向反向
op 当前 op
*args, **kwargs 该 op 的实际输入参数

返回值(CheckpointPolicy 枚举):

返回值 效果
MUST_SAVE / PREFER_SAVE 前向保留激活值,反向直接复用,不重计算
MUST_RECOMPUTE / PREFER_RECOMPUTE 反向时重新运算该 op
MUST_SWAP 前向将激活 offload 到 CPU,反向异步加载回来(必须配合 SwapManager 使用)
swap_wrapper 的 policy_fn

函数签名:

policy_fn(tensor) -> CheckpointPolicy
参数 说明
tensor autograd 引擎即将保存的 tensor

返回值(CheckpointPolicy 枚举):

返回值 效果
CheckpointPolicy.MUST_SAVE 该 tensor 保留在设备上,不 offload
CheckpointPolicy.MUST_SWAP 该 tensor 被异步 offload 到 CPU (必须配合 SwapManager 使用)

基于DispatchMode的选择重计算:

定义于hyper-parallel\hyper_parallel\platform\mindspore\activation_checkpoint\sac.py的 create_selective_checkpoint_contexts,本质上是在给重计算提供一对“前向态 / 反向态”的上下文拦截。
它返回一对 context:

_CachingMindSporeDispatchMode(...)
_CachedMindSporeDispatchMode(...)

完整调用链

checkpoint_wrapper(module, policy_fn=my_policy)
  └─ checkpoint(function, policy_fn=my_policy)           # activation_checkpoint.py:59
       context_fn = partial(create_selective_checkpoint_contexts, policy_fn)
       plat.checkpoint(function, context_fn=context_fn, use_reentrant=False)
         └─ recompute(block, use_reentrant=False, context_fn=context_fn)  # recompute.py:430
              └─ recompute_without_reentrant(block, ..., context_fn)       # recompute.py:517

forward_ctx, recompute_ctx = context_fn()

前向传播 — forward_ctx
with _CreatePlaceholderHook(state), forward_ctx:
    result = wrapper_block(*args, **kwargs)

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:

def recompute_function(*inputs):
    ...
    with recompute_ctx:  # _CachedMindSporeDispatchMode,同样拦截每个算子
        wrapper_block(*args, **kwargs)  # 重走 forward 让 autograd 引擎重新执行每个算子,从而填充那些 placeholder

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 和中间激活的处理精确地分离开来:

时间窗口 活跃 hook 拦截内容
_InputSaver.apply() 调用期间 AsyncSaveOnCpu ctx.save_for_backward(*inputs) → 由 hook 将 inputs 卸载到 CPU
_CreatePlaceholderHook 开启后 _CreatePlaceholderHook block 内中间激活 → 原重计算流程

Swap 覆盖范围:

Swap 机制统一管理 Checkpoint Activation,覆盖以下两类对象:

类型 说明
Operator Output Swap 对指定算子的输出结果进行 offload / load
Checkpoint Input Swap 对进入 checkpoint 区域的输入 Tensor 进行 offload / load

触发时机区别如下:

  • 算子输出
    • 在算子前向执行完成后触发 offload
  • Checkpoint 输入
    • 在 checkpoint 前向入口处触发 offload
    • 在反向或重计算前触发 load

Swap 管理:

系统采用四层抽象:

SwapManager          ← 全局单例,按 group_name 路由操作
    └── SwapGroup    ← 多 Storage 的统一调度
            └── Storage    ← 管理一个模块在某次前向中保存的所有 tensor 集合(save/swap 分开)
                    └── SwapTensor   ← 单个 tensor 的状态机
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 时机,实现计算与拷贝并发。

  • 正向过程中first层执行结束后立即做offload,并执行上一层的device内存清理;
  • 反向过程中second开始执行时做first层的prefetch,并执行当前层的wait_load;
  • 对于没有后继的最后一层,正向不做offload,没有前驱的第一层,不做prefetch。

输入图片说明

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)

输入图片说明

通过分组机制实现:

  • 不同模块 / stage 的 swap 隔离
  • 更精细的 prefetch 与释放控制

涉及到的对外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模式峰值显存应下降。

likedislike
DavidFFFanDavidFFFan成员
1月19日 创建了Requirement
DavidFFFanDavidFFFan成员
1月19日 关联了MindSpore/hyper-parallel Pull Request !115
DavidFFFanDavidFFFan成员
1月19日 修改了标题
DavidFFFanDavidFFFan成员
1月19日 修改了描述
DavidFFFanDavidFFFan成员
1月19日 修改了描述
此处折叠了90条消息 查看更多
songjiaqisongjiaqi成员
5月12日 修改了issue 的描述
songjiaqisongjiaqi成员
5月12日 修改了issue 的描述
fangwenyifangwenyi成员
6月23日 issue状态由 TODO 改变为 DONE
fangwenyifangwenyi成员
6月23日 关闭了 issue
MindSpore-Bot
MindSpore-Bot成员
6月23日 评论:

Notice

@ , this issue is linked to an open PR. Please merge the PR before closing this issue.

likedislike