mindspore-ai / mindspore-ai/hyper-parallel
[RFC] Torch Qwen3-30B-A3B Attention 激活内存 Swap
Nobody has claimed this yet.
- Dominant language
- Python
- Stars
- 53
- Forks
- 63
- Avg merge
- 23h 45m
- Merged PRs (30d)
- 63
Description
1. 基本信息
| 项目 | 内容 |
|---|---|
| 作者 | 待开发者认领 |
| 相关模块 | model / activation_checkpoint / trainer / pipeline_parallel(仅复用) |
| 相关 issue / PR | https://gitcode.com/mindspore/hyper-parallel/pull/1177 |
| 适用后端 | PyTorch |
| 适用模型 | Hugging Face Qwen3MoeForCausalLM,本期仅验收 Qwen3-30B-A3B |
2. 背景
Qwen3-30B-A3B 训练时,Attention 前向中由 autograd 保存、供反向使用的激活会在设备侧持续占用显存。Hyper-Parallel 已具备通用的 saved-tensor swap 能力、非 PP 的逐层 offload/prefetch 能力,以及 PP 场景下按 (stage, microbatch) 管理 swap 生命周期的调度能力,但新 Trainer 尚未提供一个只针对 Hugging Face Qwen3-30B-A3B Attention 的统一启用入口和模型 patch。
本 RFC 要解决的问题:在不重算 Attention、不修改 PP 调度流程的前提下,将 Qwen3-30B-A3B Attention 内满足条件的 saved-for-backward tensor 异步换出到 CPU pinned memory,并在反向使用前预取回设备。
完成后的成功标准:同一公共配置可同时覆盖 PP 与非 PP;关闭配置时行为不变;开启后训练 loss/梯度与基线对齐,且实测设备峰值显存下降。
3. 目标和非目标
3.1 目标
- 仅支持 PyTorch 后端的 Hugging Face
Qwen3MoeForCausalLM,验收模型为 Qwen3-30B-A3B。 - 只 swap
Qwen3MoeAttention范围内 autograd 保存用于 backward 的合格激活,不 swap DecoderLayer 的两个 RMSNorm、residual 和 MLP/MoE 激活。 - PP 场景完全复用已有 PP swap schedule,以
(stage, microbatch)为 swap group;本特性只负责注册 Attention saved tensor。 - 非 PP 场景复用
SwapManager.set_forward_prefetch_layer(),按本地 Attention 层顺序完成 D2H/H2D 调度。 - 使用实例级 patch/wrapper,不复制或重写 Hugging Face Attention forward,保持 state_dict 参数名兼容。
- 默认关闭,关闭后不增加 hook、不包装模型、不改变现有训练流程。
3.2 非目标
- 不支持 MindSpore 后端;开启时应明确报错,不做静默降级。
- 不支持
torch.compile或其他图编译;swap attention 与图编译同时开启时应在模型准备阶段明确报错。 - 不新增、不修改 PP stage 切分、microbatch 编排、1F1B/interleaved schedule 和 PP 通信逻辑。
- 不支持其他模型、Qwen3-VL/Omni、Qwen3-Next 或自定义 Qwen3 派生类;未经验收的模型不承诺可用。
- 不实现 MLP/MoE 重算,不改变已有 gradient checkpointing 语义,不支持 Attention swap 与整层 full checkpoint 同时覆盖同一模块。
- 本期不做
(stage, microbatch, layer)三级 PP 子 group;PP 下沿用已有 stage-microbatch 粒度。
4. 相关实现参考
| 来源 | 做法 | 限制 | 对本 RFC 的影响 |
|---|---|---|---|
Hyper-Parallel swap_wrapper / saved_tensors_hooks |
在 forward context 中收集 autograd 保存的 tensor | 必须存在有效 swap group | 作为 Attention patch 的核心能力直接复用 |
| Hyper-Parallel PP swap | schedule 注入 SET_GROUP、D2H、H2D、WAIT 步骤,group 为 (stage, microbatch) |
粒度是 stage-microbatch | PP 侧不新增模型 hook 调度,只注册 tensor |
SwapManager.set_forward_prefetch_layer |
相邻层 forward 后异步 D2H,backward 前预取前一层 | module backward hook 与 FSDP 组合需注意 view/in-place | 非 PP 侧复用;必要时复用现有 tensor backward hook 方案 |
| Hugging Face Qwen3-MoE | Qwen3MoeDecoderLayer.self_attn 边界清晰 |
Transformers 升级可能调整类名或层级 | 只包 self_attn,避免重写 forward,降低版本耦合 |
5. 对外接口
5.1 接口定义
在新 Trainer 的公共配置中增加一个开关:
activation_swap: Literal["none", "attention"] = "none"
YAML 示例:
activation_swap: attention
| 配置项 | 类型 | 默认值 | 是否必填 | 含义 | 合法范围 | 错误处理 |
|---|---|---|---|---|---|---|
activation_swap |
str |
none |
否 | 激活 swap 模式 | none / attention |
非法值在配置解析阶段报错 |
本期只保留一个公共配置。tensor 大小阈值、group copy 等先使用实现侧常量和已有默认策略,不在本期扩展更多用户接口;后续如有多模型实测需求,再单独开放调优参数。
5.2 使用示例
model:
# HF Qwen3-30B-A3B 配置
activation_swap: attention
activation_swap=none:完全保持现有行为。
activation_swap=attention + pp_size>1:使用已有 PP swap schedule。
activation_swap=attention + pp_size=1:安装 Attention 逐层调度 hook。
5.3 参数校验
开启 attention 时必须满足:
- 当前平台为 PyTorch;否则抛出
NotImplementedError。 - 未开启图编译;否则抛出
ValueError,错误信息说明两者不兼容。 - 模型为 HF
Qwen3MoeForCausalLM,DecoderLayer 中存在预期的Qwen3MoeAttention;否则抛出ValueError。 - 不允许 HF native full gradient checkpointing 再覆盖同一 DecoderLayer;发现冲突时明确报错或按最终实现统一关闭,不能静默双重包装。
6. 方案设计
6.1 总体流程
flowchart TD
A[读取 activation_swap] --> B{是否为 attention}
B -- 否 --> C[保持现有流程]
B -- 是 --> D[校验 Torch/模型/无图编译]
D --> E[定位当前 rank 的 Qwen3MoeDecoderLayer]
E --> F[仅使用 swap_wrapper 包装 self_attn]
F --> G{pp_size > 1?}
G -- 是 --> H[不注册逐层调度 hook]
H --> I[复用 PP stage-microbatch swap schedule]
G -- 否 --> J[按 Attention 顺序注册逐层 offload/prefetch]
I --> K[执行训练]
J --> K
6.2 模块边界
Qwen3MoeDecoderLayer
├── input_layernorm 不在 swap context
├── self_attn 使用 swap_wrapper
│ ├── q_proj / k_proj / v_proj
│ ├── q_norm / k_norm / RoPE
│ ├── fused/SDPA attention core
│ └── o_proj
├── residual add 不在 swap context
├── post_attention_layernorm 不在 swap context
├── mlp / sparse MoE 不在 swap context
└── residual add 不在 swap context
swap_wrapper 捕获的是 Attention 内部由 autograd 保存供 backward 使用的 tensor,不是无条件复制所有中间结果。参数、参数 view、空 storage、小 tensor、无梯度 tensor以及不安全的共享 storage view应由基础检查或模型专用 policy 保留在设备侧。
建议的内部 policy:
def qwen3_attention_swap_policy(tensor):
if not tensor.requires_grad:
return CheckpointPolicy.MUST_SAVE
if tensor.numel() * tensor.element_size() < MIN_SWAP_TENSOR_BYTES:
return CheckpointPolicy.MUST_SAVE
if tensor.untyped_storage().size() != tensor.numel() * tensor.element_size():
return CheckpointPolicy.MUST_SAVE
return CheckpointPolicy.MUST_SWAP
MIN_SWAP_TENSOR_BYTES 使用实现侧常量,并通过 Qwen3-30B-A3B 实测确定初始值。
6.3 模型 patch
采用实例级包装,不做全局 class monkey patch:
def patch_qwen3_30b_a3b_attention_swap(model, config):
validate_backend_model_and_compile(model, config)
local_layers = find_local_qwen3_moe_layers(model)
wrapped_attentions = []
for layer in local_layers:
if already_wrapped(layer.self_attn):
continue
layer.self_attn = swap_wrapper(
layer.self_attn,
policy_fn=qwen3_attention_swap_policy,
group_swap=True,
)
wrapped_attentions.append(layer.self_attn)
if config.accelerator.pp_size == 1:
setup_non_pp_attention_schedule(wrapped_attentions)
mark_patch_installed(model)
return model
patch 应在 PP 完成本地 stage 划分之后、FSDP/sharding 包装之前执行;只扫描和包装当前 rank 实际持有的 DecoderLayer。必须提供幂等标记,重复调用不能重复 wrapper 或重复注册 hook。
6.4 PP 场景时序
sequenceDiagram
participant S as PP Scheduler
participant M as Local PP Stage
participant A as Patched Attention
participant W as SwapManager
S->>W: set group(stage, microbatch)
S->>M: forward(microbatch)
M->>A: attention forward
A->>W: register saved tensors
M-->>S: stage output
S->>W: launch/wait D2H
Note over S,W: 复用原有 PP 调度和间隔判断
S->>W: launch/wait H2D
S->>M: backward(microbatch)
本特性不得在 PP 场景额外注册非 PP 的逐层 group hook,否则会覆盖 scheduler 设置的 current group。
6.5 非 PP 场景时序
对包装后的本地 Attention 顺序调用:
for current_attn, next_attn in pairwise(wrapped_attentions):
SwapManager().set_forward_prefetch_layer(current_attn, next_attn)
sequenceDiagram
participant A0 as Attention N
participant W as SwapManager
participant A1 as Attention N+1
A0->>W: set current layer group
A0->>A0: forward/save tensors
A0->>W: launch D2H
A1->>W: wait previous D2H
A1->>A1: forward/save tensors
Note over A0,A1: 中间的 MoE 和后续计算可隐藏 D2H
A1->>W: backward pre-hook,launch load Attention N
A1->>A1: backward
A0->>W: wait load
A0->>A0: backward
如 FSDP2 下 module full-backward hook 与 view/in-place 操作冲突,应复用仓库已有的 tensor backward hook 方案:在第一个 requires-grad output 上注册 backward-start hook,在第一个 requires-grad input 上注册 backward-end hook。本 RFC 不新设计另一套 swap manager。
6.6 代码改动点
| 模块 | 改动内容 | 是否影响已有行为 |
|---|---|---|
trainer/config |
增加 activation_swap 公共配置及合法值校验 |
默认 none 时不影响 |
| Torch 模型 patch | 增加 Qwen3-30B-A3B Attention 识别、包装和幂等处理 | 只在显式开启时影响目标模型 |
| activation checkpoint/swap 集成 | 增加 Attention policy;必要时将现有 tensor backward hook 提升为公共复用函数 | 默认不影响 |
| PP scheduler | 不修改;只消费已有 PP swap 能力 | 否 |
| MindSpore | 不实现 | 否 |
| graph compile | 不实现;增加冲突校验 | 仅开启本特性时校验 |
6.7 方案取舍
| 方案 | 优点 | 缺点 | 是否选择 | 原因 |
|---|---|---|---|---|
实例级 swap_wrapper(self_attn) |
改动小,不复制 HF forward,参数名可兼容,PP/非 PP 共用 | 依赖 HF 模块边界 | 是 | 满足单模型、只 swap Attention 的范围 |
重写 Qwen3MoeAttention.forward |
可精确包 fused attention core | Transformers 版本耦合强,维护成本高 | 否 | 本期没有精确到算子级的必要 |
| 包装整个 DecoderLayer | 接入简单 | 会捕获 norm/MoE 激活,超出需求 | 否 | 与“只 swap Attention”冲突 |
| 为 PP 新增逐层子 group | 可更早释放 stage 内激活 | 需要改 PP 调度和 backward 编排 | 否 | PP 已有流程,本期明确不改 |
7. 组件依赖
| 依赖组件 | 强依赖 / 弱依赖 | 当前状态 | 未 ready 时本期能力 |
|---|---|---|---|
| PyTorch saved tensor hooks | 强依赖 | 已有 | 无法提供本特性 |
swap_wrapper / SwapManager |
强依赖 | 已有 | 无法提供本特性 |
| PP swap schedule | PP 场景强依赖 | 已适配 | 非 PP 仍可工作 |
| FSDP2 | 弱依赖 | 使用现有能力 | 可先完成无 FSDP Level0;正式验收需覆盖目标组合 |
| MindSpore | 不涉及 | 本期不支持 | 明确报错 |
| 图编译 | 不涉及 | 本期不支持 | 明确报错 |
完整能力需要:Trainer 在模型准备阶段调用 patch;PP 场景构造 schedule 时已启用现有 swap;非 PP 场景可安装逐层 Attention hook。
本期最小可交付能力:Torch + HF Qwen3-30B-A3B 在 PP/非 PP 两种场景下可通过同一配置开启 Attention swap。
8. 约束与兼容性
| 类型 | 内容 |
|---|---|
| 不支持项 | MindSpore、图编译、非 Qwen3-30B-A3B 模型、Qwen3-VL/Omni、PP 逐层子 group、MLP/MoE 重算 |
| 默认兼容性 | activation_swap=none 时不包装、不注册 hook,行为与当前版本一致 |
| checkpoint | state_dict key 必须与未包装模型兼容;save/load 后可继续训练 |
| PP 差异 | PP 按 (stage, microbatch) 调度;非 PP 按 Attention 层调度 |
| 显存收益 | 取决于序列长度、Attention 实现、PP microbatch 数和可 swap saved tensor 数量;验收要求峰值显存低于基线,不预设未经实测的固定比例 |
| 性能代价 | 增加 D2H/H2D;通过异步 copy、pinned pool、group copy 和计算重叠降低影响,需输出实测数据 |
| Transformers 版本 | 依赖 Qwen3MoeForCausalLM -> model.layers[*].self_attn 边界;结构不匹配时快速失败,不能静默漏 patch |
9. 验证设计
9.1 用例分层
| 用例级别 | 建议数量 | 覆盖内容 | 通过标准 |
|---|---|---|---|
| UT | 8~10 | 配置默认值/非法值、后端校验、图编译冲突、模型识别、只包装 Attention、幂等、policy、PP/非 PP 分流 | 全部通过;默认配置无副作用 |
| Level0 | 2 | tiny Qwen3Moe 单进程基线与非 PP swap;最小 PP smoke | loss/梯度对齐,无残留 group,无非法 storage 状态 |
| Level1 | 2~4 | Qwen3-30B-A3B 非 PP、PP;按实际训练方案补 FSDP2/TP/EP 组合 | 连续训练稳定,峰值显存下降,性能数据可解释 |
9.2 核心正确性验证
- 固定随机种子、输入、dtype 和优化器,对比关闭/开启 swap 的 forward loss 和 parameter gradients。
- Swap 是纯数据搬运、无重算,FP32 单测期望 loss 一致,梯度使用
torch.testing.assert_close;BF16 设备测试沿用项目既有精度阈值。 - 通过计数或测试 spy 证明实际发生了 Attention tensor 注册、D2H、H2D,而不是只完成包装但未进入有效 group。
- 检查 DecoderLayer 的两个 RMSNorm 和 MLP/MoE 没有加入 swap storage。
- 重复 prepare 不得增加 wrapper 层数或 hook 数量。
- 训练 step 结束和异常路径后,swap group/storage/stream 状态可被清理,下一 step 可继续运行。
9.3 交互验证
| 组合 | 是否验证 | 通过标准 |
|---|---|---|
| Attention swap + 非 PP | 是 | 按层 D2H/H2D,loss/梯度对齐 |
| Attention swap + PP | 是 | 使用原 PP schedule,不安装非 PP hook,loss/梯度对齐 |
| Attention swap + FSDP2 | 按目标训练配置验证 | 无 backward-hook view/in-place 冲突,参数和 checkpoint 正常 |
| Attention swap + TP/EP | 若 Qwen3-30B-A3B 验收配置启用则验证 | 训练稳定,loss 符合项目阈值 |
| Attention swap + graph compile | 否 | prepare 阶段按预期报错 |
| Attention swap + MindSpore | 否 | prepare 阶段按预期报错 |
| Attention swap + full layer checkpoint | 不支持 | 配置或 prepare 阶段按预期报错,避免重复覆盖 |
9.4 性能和显存验证
| 场景 | 基线 | 开启特性 | 指标 | 通过标准 |
|---|---|---|---|---|
| 非 PP Qwen3-30B-A3B | activation_swap=none |
attention |
peak device memory、step time、D2H/H2D bytes/time | 峰值显存下降;报告吞吐变化和拷贝带宽 |
| PP Qwen3-30B-A3B | 原 PP、swap 关闭 | 原 PP schedule + Attention 注册 | 各 rank peak memory、step time、swap group 数 | 至少目标 rank 峰值显存下降;PP 调度无新增修改和死锁 |
性能不预设未经实测的固定收益比例。PR 验收材料必须同时给出:模型配置、序列长度、global/micro batch、并行策略、rank 数、Attention 实现、基线/开启后的显存和吞吐。
Qwen3-30B-A3B seq_length:1024
| 配置 | peak_mem | 显存优化率 | step_time | 性能劣化率 |
|---|---|---|---|---|
none |
30.14G | / | 5.3123s | / |
attention |
27.17G | 10% | 5.5957s | 5.33% |
10. 实现计划与工作量
| PR | 内容 | 依赖 | 验证 | AI 辅助后预计 |
|---|---|---|---|---|
| PR1 | 公共配置、冲突校验、Qwen3 Attention patch、policy、幂等 | 现有 swap_wrapper |
UT | 0.5~0.75 人天 |
| PR2 | 非 PP Attention 链式调度;PP 复用路径接线 | PR1、现有 SwapManager 和 PP schedule |
UT + Level0 | 0.25~0.5 人天 |
| PR3 | Qwen3-30B-A3B PP/非 PP 精度、显存和性能验证,补问题修复 | PR2、设备环境 | Level1 | 0.5~1.25 人天 |
总计:使用 AI 辅助约 1.25~2.5 人天。其中核心代码和 UT 约 0.75~1.25 人天,真实设备验证约 0.5~1.25 人天。设备排队和环境问题不计入纯开发工时。
11. 验收 Checklist
- 默认配置为
none,现有训练行为不变。 - 只支持 Torch + HF Qwen3-30B-A3B;其他模型快速失败。
- 只包装
self_attn,DecoderLayer norm、residual、MLP/MoE 未进入 swap。 - PP 复用原有 PP swap schedule,未注册非 PP layer hook,未修改 PP 编排。
- 非 PP 按 Attention 层完成异步 offload/prefetch。
- MindSpore 和图编译场景给出清晰的不支持错误。
- patch 幂等,state_dict key 兼容,checkpoint 可恢复训练。
- loss/梯度满足项目精度标准。
- 实测发生 D2H/H2D,且目标场景峰值设备显存下降。
- PR 提供完整的 Qwen3-30B-A3B 显存与吞吐对比数据。
schema_version: 1
source: gitcode
gitcode_repo: mindspore/hyper-parallel
gitcode_issue: 323
source_url: https://gitcode.com/mindspore/hyper-parallel/issues/323
Contributor guide
No contributing guide indexed for this repository
First steps
- Read the whole issue, then the project's contributing guide.
- Comment on the issue to say you are picking it up — it saves two people doing the same work.
- Fork the repository and make your change on a branch.
- Open a pull request that references the issue number.
Research direction
Start by reading the existing swap_wrapper/saved-tensor hooks, SwapManager, and PP swap schedule, then trace the trainer/config model-preparation entry point. Locate the Hugging Face Qwen3MoeForCausalLM self_attn boundary described in the RFC and review the proposed validation and idempotence requirements. Done means one activation_swap configuration supports PP and non-PP attention swapping with the listed rejection checks, tests, loss/gradient alignment, and measured memory results.
Written by the indexing model from the issue text.
Assessment
- Tech stack
- python
- Domain
- backend, machine-learning, performance
- Issue type
- Feature
- Difficulty
- 4/5
- Estimated time
- 3-5 days
- Activity status
- Active
- Clarity
- Mostly clear
- Newbie friendliness
- 42/100