mindspore-ai / mindspore-ai/hyper-parallel

[RFC] Torch Qwen3-30B-A3B Attention 激活内存 Swap

Open
#192 2 comments 0 reactions 0 assignees View on GitHub

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 目标
  1. 仅支持 PyTorch 后端的 Hugging Face Qwen3MoeForCausalLM,验收模型为 Qwen3-30B-A3B。
  2. 只 swap Qwen3MoeAttention 范围内 autograd 保存用于 backward 的合格激活,不 swap DecoderLayer 的两个 RMSNorm、residual 和 MLP/MoE 激活。
  3. PP 场景完全复用已有 PP swap schedule,以 (stage, microbatch) 为 swap group;本特性只负责注册 Attention saved tensor。
  4. 非 PP 场景复用 SwapManager.set_forward_prefetch_layer(),按本地 Attention 层顺序完成 D2H/H2D 调度。
  5. 使用实例级 patch/wrapper,不复制或重写 Hugging Face Attention forward,保持 state_dict 参数名兼容。
  6. 默认关闭,关闭后不增加 hook、不包装模型、不改变现有训练流程。
3.2 非目标
  1. 不支持 MindSpore 后端;开启时应明确报错,不做静默降级。
  2. 不支持 torch.compile 或其他图编译;swap attention 与图编译同时开启时应在模型准备阶段明确报错。
  3. 不新增、不修改 PP stage 切分、microbatch 编排、1F1B/interleaved schedule 和 PP 通信逻辑。
  4. 不支持其他模型、Qwen3-VL/Omni、Qwen3-Next 或自定义 Qwen3 派生类;未经验收的模型不承诺可用。
  5. 不实现 MLP/MoE 重算,不改变已有 gradient checkpointing 语义,不支持 Attention swap 与整层 full checkpoint 同时覆盖同一模块。
  6. 本期不做 (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 时必须满足:

  1. 当前平台为 PyTorch;否则抛出 NotImplementedError
  2. 未开启图编译;否则抛出 ValueError,错误信息说明两者不兼容。
  3. 模型为 HF Qwen3MoeForCausalLM,DecoderLayer 中存在预期的 Qwen3MoeAttention;否则抛出 ValueError
  4. 不允许 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 核心正确性验证
  1. 固定随机种子、输入、dtype 和优化器,对比关闭/开启 swap 的 forward loss 和 parameter gradients。
  2. Swap 是纯数据搬运、无重算,FP32 单测期望 loss 一致,梯度使用 torch.testing.assert_close;BF16 设备测试沿用项目既有精度阈值。
  3. 通过计数或测试 spy 证明实际发生了 Attention tensor 注册、D2H、H2D,而不是只完成包装但未进入有效 group。
  4. 检查 DecoderLayer 的两个 RMSNorm 和 MLP/MoE 没有加入 swap storage。
  5. 重复 prepare 不得增加 wrapper 层数或 hook 数量。
  6. 训练 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

  1. Read the whole issue, then the project's contributing guide.
  2. Comment on the issue to say you are picking it up — it saves two people doing the same work.
  3. Fork the repository and make your change on a branch.
  4. 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

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.