已开启
【需求】Pipeline场景支持SWAP #174
DavidFFFan创建于 5月29日
5月29日 修改了issue 的描述
5月29日 修改了issue 的描述
5月30日 关联了pull request:pp support swap
5月30日 将 DavidFFFan 设为负责人
7月29日 修改了issue 的描述
7月29日 修改了issue 的描述
7月30日 修改标题为 “【需求】Pipeline场景支持SWAP”,原标题为“Pipeline场景支持SWAP”
7月30日 修改标题为 “Pipeline场景支持SWAP”,原标题为“【需求】Pipeline场景支持SWAP”
7月30日 修改了issue 的描述
7月30日 修改了issue 的描述
7月30日 修改标题为 “【需求】Pipeline场景支持SWAP”,原标题为“Pipeline场景支持SWAP”
8月6日 将 Meng107 设为负责人
8月6日 issue状态由 TODO 改变为 ACCEPTED
8月10日 修改了issue 的描述
hui-zhang940
8月12日 评论:
8月12日 评论:
如何确认某个 tensor 被正确 swap?是否有验证方式?


hui-zhang940
8月12日 评论:
8月12日 评论:
swap_tensor_wrapper没转测过,是新增的接口么


hui-zhang940
8月12日 评论:
8月12日 评论:


hui-zhang940
8月12日 评论:
8月12日 评论:


8月23日 修改了issue 的描述
Pipeline 场景支持 Activation Swap
1. 基本信息
pipeline_parallel/activation_checkpoint/fully_shard2. 背景
Pipeline Parallel(PP)训练会为多个 micro-batch 保留前向激活,直到对应反向到来。GPipe、Interleaved 1F1B、real overlap 等调度中,部分
(stage_index, micro_index)的 FWD 与 BWD 之间存在较长空窗,激活持续驻留 device 会抬高峰值显存。HyperParallel 已有 activation swap 基础能力,但基础 swap 本身不知道 PP 的 micro-batch、virtual stage、FSDP MetaStep、P2P container 和 overlap 线程边界,无法安全选择 D2H/H2D 时机。因此本特性不新增另一套 swap 接口,而是在用户通过
swap_wrapper或swap_tensor_wrapper使能 swap 的基础上,由 PP scheduler 负责搬运调度。3. 目标和非目标
3.1 目标
overlap_b_f和 dx/dw 拆分场景的 PP swap 调度。(stage_index, micro_index)建立独立 swap group,支持同一 rank 上的多个 virtual stage。swap_wrapper,指定 tensor 使用swap_tensor_wrapper;PP 只增加调度能力。schedule.run()使用独立 generation,正常连续运行时 group 不串轮次。3.2 非目标
MIN_SWAP_GAP = 4。4. 相关实现参考
SwapTensor、Storage、SwapGroup、SwapManager管理 saved tensor、D2H/H2D、buffer 和 group 生命周期swap_wrapperswap_tensor_wrapper5. 对外接口
5.1 接口定义
PP swap 由“激活选择”和“PP 调度”两部分组成,二者缺一不可:
from hyper_parallel.core.activation_checkpoint import swap_tensor_wrapper, swap_wrapper # 方式一:模块级。包裹模块后,模块 forward 中需要保存的激活按 policy_fn 参与 swap。 wrapped_module = swap_wrapper( module, policy_fn=None, group_swap=False, ) # 方式二:tensor 级。在 forward/construct 内显式注册指定 tensor。 target = swap_tensor_wrapper( target, tag=None, group_swap=False, ) # PP 侧只负责调度。swap=True 不会替代上述 wrapper,也不会自动选择激活。 schedule = ScheduleGPipe( stage, micro_batch_num=8, swap=True, )swap_wrapper.moduleCell/nn.Module/ callableswap_wrapper.policy_fnNoneNoneswap_wrapper.group_swapboolFalseTrue/Falseswap_tensor_wrapper.targetswap_tensor_wrapper.tagstr/NoneNoneswap_tensor_wrapper.group_swapboolFalseTrue/FalseSchedule*.swapboolFalseTrue/False5.2 使用示例
5.2.1 模块级 swap
from hyper_parallel import PipelineStage, ScheduleInterleaved1F1B from hyper_parallel.core.activation_checkpoint import swap_wrapper # 先圈定需要 swap 的层;赋值回模型后 wrapper 才会生效。 for index, layer in enumerate(stage_model.layers): stage_model.layers[index] = swap_wrapper(layer, group_swap=True) stage = PipelineStage(stage_model, stage_index, stage_num=pp_size) # swap=True 仅开启 PP 搬运调度。 schedule = ScheduleInterleaved1F1B( stages=[stage], micro_batch_num=8, swap=True, ) losses = schedule.run(*inputs)5.2.2 指定 tensor swap
from hyper_parallel import PipelineStage, ScheduleGPipe from hyper_parallel.core.activation_checkpoint import swap_tensor_wrapper class TransformerBlock(nn.Cell): def construct(self, x): attn_out = self.attn(x) attn_out = swap_tensor_wrapper(attn_out, tag="attn_out", group_swap=True) x = self.norm1(x + attn_out) return x stage = PipelineStage(stage_model, stage_index, stage_num=pp_size) schedule = ScheduleGPipe(stage, micro_batch_num=8, swap=True) losses = schedule.run(*inputs)swap_tensor_wrapper必须在 stage 的 forward/construct 路径内调用。PP runtime 会在 eligible FWD leaf 外建立 swap group context;如果该 chunk 没有有效 swap window,则不会创建物理 group,也不会执行搬运。5.2.3 与 FSDP 组合
from hyper_parallel import PipelineStage, ScheduleGPipe, fully_shard from hyper_parallel.core.activation_checkpoint import swap_wrapper # 推荐先对 layer/子层应用 swap wrapper,再对整个 PP stage 应用 fully_shard。 for index, layer in enumerate(stage_model.layers): stage_model.layers[index] = swap_wrapper(layer) fully_shard(stage_model, mesh=dp_mesh) stage = PipelineStage(stage_model, stage_index, stage_num=pp_size) schedule = ScheduleGPipe(stage, micro_batch_num=8, swap=True)5.3 接口说明
6. 方案设计
6.1 总体流程
flowchart TD A["用户应用 swap_wrapper 或 swap_tensor_wrapper"] --> B["构建 PP Schedule,swap=True"] B --> C["构建 compute / P2P order"] C --> D["注入本 rank FSDP actions"] D --> E["重写 P2P transport"] E --> F["分析 FWD/BWD leaf 与 gap"] F --> G{"gap >= 4?"} G -- "否" --> H["chunk 保持 device resident"] G -- "是" --> I["注入 4 个 swap MetaStep"] I --> J["run-scoped session 建立 chunk group"] J --> K["FWD 收集激活"] K --> L["D2H offload / H2D load"] L --> M["SWAP_WAIT_LOAD 后执行 BWD consumer"] M --> N["run close 回收 group"]build_exec_order()的顺序固定为:swap 最后注入,才能看到最终 FSDP lookahead 和 P2P container;swap MetaStep 不进入
BATCH_SEND_RECV.sub_steps。6.2 架构设计
flowchart LR subgraph User["用户侧:选择激活"] SW["swap_wrapper"] STW["swap_tensor_wrapper"] end subgraph PP["PP 调度侧:选择时机"] Planner["inject_pipeline_swap_steps"] Session["PipelineSwapSession"] Executor["MetaStep executor"] end subgraph Base["基础 Swap Runtime"] Manager["SwapManager"] Group["SwapGroup / Storage / SwapTensor"] Copy["D2H / H2D copy stream"] end SW --> Group STW --> Group Planner --> Executor Executor --> Session Session --> Manager Manager --> Group Group --> Copy职责边界:
swap_wrapper、swap_tensor_wrapperSwapTensor、Storage、SwapGroup、SwapManagerinject_pipeline_swap_steps()PipelineSwapSession、MetaStep executor6.3 单个 chunk 时序与 4 个 MetaStep
当前实现使用 4 个显式 MetaStep,不再是 3 个:
sequenceDiagram participant F as FWD leaf participant S as PP scheduler participant C as Copy stream participant B as BWD/BWD_INPUT F->>F: 在 chunk group context 中收集激活 S->>C: SWAP_LAUNCH_OFFLOAD S->>S: SWAP_WAIT_OFFLOAD Note over S: 建立 event 依赖并释放 device storage S->>C: SWAP_LAUNCH_LOAD S->>S: SWAP_WAIT_LOAD Note over S: 在 backward consumer container 前建立 H2D 依赖 S->>B: 执行 BWD 或 BWD_INPUTDEVICESWAP_LAUNCH_OFFLOADDEVICE → D2HSWAP_WAIT_OFFLOADD2H → HOSTSWAP_LAUNCH_LOADHOST → H2DSWAP_WAIT_LOADH2D → DEVICEBWD_INPUTleafDEVICEwait_load()与 run close逻辑 chunk key 为
(stage_index, micro_index),物理 group name 为:6.4 Planner 关键逻辑
6.4.1 Leaf 提取
FWD、BWD、BWD_INPUT、BWD_WEIGHT是 compute leaf。OVERLAP_B_F/OVERLAP_F_B展开sub_steps用于配对,但原 composite container 保持不变。FWD → BWD配对;dx/dw 路径按FWD → BWD_INPUT配对。BWD_WEIGHT不是 activation first consumer,不参与 load 配对。6.4.2 Eligibility
MIN_SWAP_GAP是收益门槛,不是正确性同步条件:6.4.3 四个静态 anchor
SWAP_LAUNCH_OFFLOADFSDP_RESHARD前SWAP_WAIT_OFFLOADSWAP_LAUNCH_LOADFSDP_UNSHARD前SWAP_WAIT_LOADBWD_INPUTtop-level consumer container 前6.5 Run-scoped session 与线程边界
每次
schedule.run()根据当前 rank order 中的SWAP_LAUNCH_OFFLOAD建立 eligible key 集合。未入选 chunk 的 FWD 使用nullcontext,不创建物理 group,也不执行 load wait。SwapManager使用ContextVar保存当前 group,支持嵌套恢复,并隔离不同 Python execution context:with session.group_context(fwd_step): output = stage.forward_one_chunk(...) session.protect_aliases(fwd_step, output)real overlap callback 必须调用统一的
execute_fwd_leaf()/execute_bwd_leaf(),不能绕过 leaf API 直接调用 stage。CommComputeOverlap.run()在主线程执行 FWD、daemon worker 执行 BWD,并在返回前 join。双 Python 线程只表示 host 可以并发下发,device 是否重叠仍取决于 stream、依赖和硬件资源。6.6 FSDP、P2P 与 dx/dw
FSDP 目标顺序:
FSDP_UNSHARD → FWD(collection)FWD → SWAP_LAUNCH_OFFLOAD → FSDP_RESHARDSWAP_LAUNCH_LOAD → FSDP_UNSHARD → SWAP_WAIT_LOAD → BWD如果目标 stage 参数已经保持 unsharded,没有可覆盖的 lookahead,则 H2D 延迟到 BWD container 前,避免 activation 过早回到 device。
P2P transport 关系:
plainbatchboundarydx/dw 拆分时,
BWD_INPUT是 activation first consumer,BWD_WEIGHT只使用 dx 阶段保存的中间状态。因此只在BWD_INPUT前插入SWAP_WAIT_LOAD,BWD_WEIGHT不重复 load 或 release;stage 0 的 backward 保持统一BWD。6.7 Stream、内存与 alias 生命周期
D2H:
H2D:
正常
wait_offload()/wait_load()使用 event 建立 device stream 依赖,不新增 hostevent.synchronize();异常 teardown 才会对 group 的 in-flight event 做定点同步。当前 alias 保护包括:
set_()alias 到恢复后的连续 device buffer。当前不做跨 group 自动 pointer 扫描;通用跨 chunk storage 共享需要后续引入明确的 ownership/generation 模型。
6.8 调度图
调度图使用 scheduler 生成的 finalized order,并基于跨 rank 数据依赖建立全局逻辑时间轴。同一列表示全局逻辑 compute slot,不代表实测 wall-clock;copy 条形表示静态 launch→wait 覆盖窗口。
C·F表示 FWD leaf 在 collection context 中执行。W·B/W·dx表示对应 consumer container 前执行SWAP_WAIT_LOAD。offload从SWAP_LAUNCH_OFFLOAD延伸到SWAP_WAIT_OFFLOAD。load从SWAP_LAUNCH_LOAD延伸到SWAP_WAIT_LOAD。6.8.1 GPipe + swap
GPipe 先执行全部 FWD,再执行全部 BWD,每个 micro-batch 通常有较长 activation 空窗。
6.8.2 VPP / Interleaved 1F1B + swap
每个物理 rank 持有多个 virtual stage。group key 使用真实
stage_index,同一 micro id 在不同 virtual stage 之间不会混组。planner 对每个(stage, micro)独立判定,没有足够窗口的 chunk 保持 device resident。6.8.3 real overlap + swap
OVERLAP_B_F(BWD_i, FWD_j)是一个 top-level composite。主线程执行 FWD leaf,daemon worker 执行 BWD leaf,同一 slot 的 main/worker 属于同一个 composite,返回前 join。6.8.4 real overlap + dx/dw + swap
6.9 代码改动点
core/activation_checkpointcore/pipeline_parallel/pipeline_swap.pyswap=True时生效core/pipeline_parallel/scheduler.pyswap=False,不影响已有调度platform/mindspore/activation_checkpointswap_wrapper/swap_tensor_wrapper后端实现tests/ut/core/pipeline_paralleltests/mindspore/st/pipeline_parallel6.10 方案取舍
主要代价是 scheduler order 增加 4 类控制 MetaStep,并需要维护 group generation、alias 保护和异常清理。
7. 组件依赖
SwapManager、wrapper)8. 约束与兼容性
8.1 支持矩阵
8.2 兼容性与收益
swap=False为默认值;不应用 wrapper 时swap=True只有调度、没有可搬运激活;非 PP swap 接口不变8.3 正确性不变量
BWD_INPUT消费前建立 stream 依赖。resize_(0)。BWD_WEIGHT不重复执行 activation load/release。9. 验证设计
9.1 用例分层
9.2 交互验证
LAUNCH_LOAD → FSDP_UNSHARD → WAIT_LOAD → BWDBWD_INPUT前 wait;BWD_WEIGHT不重复 load9.3 性能 / 显存验证
real-overlap 两个 E2E 使用默认
auto → batchtransport,不覆盖 FSDP + real overlap 或 boundary + swap。10. 实现计划