已开启
[RFC] Torch Qwen3-30B-A3B Attention 激活内存 Swap #323
DavidFFFan创建于  8月10日
DavidFFFan
DavidFFFan成员
8月10日 创建

1. 基本信息

项目 内容
作者 待开发者认领
相关模块 distributed/attention_swap.py、models/_transformers/model_builder.py、core/activation_checkpoint、trainer
相关 issue / PR https://gitcode.com/mindspore/hyper-parallel/pull/1177
适用后端 PyTorch
适用模型 本期仅验收 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,并在反向使用前预取回设备。

3. 目标和非目标

3.1 目标

  1. 仅支持 PyTorch 后端的 Hugging Face Qwen3MoeForCausalLM,验收模型为 Qwen3-30B-A3B。
  2. 只 swap Qwen3MoeAttention 范围内 autograd 保存用于 backward 的激活,不 swap DecoderLayer 的两个 RMSNorm、residual 和 MLP/MoE 激活。
  3. 只支持非 PP 场景,通过 SwapManager.set_forward_prefetch_layer() 按 Attention 扫描顺序建立非 PP 的异步
    D2H/H2D 调度。
  4. 使用 swap_wrapper。
  5. 默认关闭 activation_swap=none。

3.2 非目标

  1. 不支持 MindSpore 后端;开启时应明确报错,不做静默降级。
  2. 不支持 torch.compile 或其他图编译;swap attention 与图编译同时开启时应明确报错。
  3. 不支持PP场景。
  4. 未经验收的模型不承诺可用,不识别属性名为 attn、attention 或其他名称的 Attention 模块;模型没有 self_attn
    子模块时会快速失败。
  5. 目前不支持 policy 可配(不提供 tensor 大小阈值、swap group、CPU pool 等用户可配)。

4. 相关实现参考

来源 做法 限制 对本 RFC 的影响
Hyper-Parallel swap_wrapper / saved_tensors_hooks 在 forward context 中收集 autograd 保存的 tensor 必须存在有效 swap group 作为 Attention patch 的核心能力直接复用
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 swap 等先使用实现侧常量和已有默认策略,不在本期扩展更多用户接口;后续如有多模型实测需求,再单独开放调优参数。

5.2 参数校验

启用 activation_swap=attention 时,validate_attention_swap() 执行以下校验:

if enable_compile:
    raise ValueError(
        "activation_swap='attention' is incompatible with torch.compile"
    )
if activation_checkpoint not in (None, "off"):
    raise ValueError(
        "activation_swap='attention' is incompatible with activation checkpointing"
    )
if pp_size != 1:
    raise ValueError(
        "activation_swap='attention' does not support pipeline parallelism"
    )
组合 当前行为
activation_swap=none 支持,保持默认行为
attention + pp_size=1 + checkpoint off + compile off 支持
attention + full/selective checkpoint 不支持,构建阶段抛出 ValueError
attention + torch.compile 不支持,构建阶段抛出 ValueError
Qwen3-MoE 且包含 self_attn 满足结构识别条件
其他模型且包含 self_attn 会按相同通用逻辑包装,但需单独验证正确性和收益
模型不包含 self_attn 不支持,抛出 ValueError

6. 方案设计

6.1 Attention 目标发现

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

从根模型开始递归调用
named_children():

  1. 子模块注册名等于 self_attn 时,将其记录为目标,并停止继续递归该目标内部;
  2. 其他子模块继续深度优先扫描;
  3. 使用模块对象的 id() 对共享模块去重;
  4. 记录同一个目标的已发现父模块和属性名,随后把这些引用替换为同一个 wrapper;
  5. 目标顺序为模块注册顺序下的深度优先首次发现顺序;
  6. 没有发现任何目标时抛出:
activation_swap='attention' found no self_attn modules in <ModelType>;
expected attention modules registered as self_attn

6.2 Tensor swap policy

当前固定策略等价于:

MIN_SWAP_TENSOR_BYTES = 1024 * 1024


def attention_swap_policy(tensor):
    if not tensor.requires_grad or tensor.dim() < 2:
        return CheckpointPolicy.MUST_SAVE

    storage_bytes = tensor.untyped_storage().size()
    tensor_bytes = tensor.numel() * tensor.element_size()
    if storage_bytes != tensor_bytes or tensor_bytes < MIN_SWAP_TENSOR_BYTES:
        return CheckpointPolicy.MUST_SAVE

    return CheckpointPolicy.MUST_SWAP

因此,下列 saved tensor 保留在设备侧:

  • 不需要梯度的 tensor;
  • 0 维或 1 维 tensor;
  • 小于 1 MiB 的 tensor;
  • 底层 storage 大小与 tensor 逻辑字节数不一致的 tensor。

6.3 模块包装

每个唯一 Attention 目标使用以下方式包装:

wrapped_attention = swap_wrapper(
    attention,
    policy_fn=attention_swap_policy,
    group_swap=True,
)

6.4 非 PP 调度

包装完成后,代码按唯一 Attention 目标的扫描顺序连接相邻模块:

swap_manager = SwapManager()
for current_attention, next_attention in zip(
    wrapped_attentions,
    wrapped_attentions[1:],
):
    swap_manager.set_forward_prefetch_layer(
        current_attention,
        next_attention,
    )

典型时序如下:

sequenceDiagram
    participant A0 as Attention N
    participant W as SwapManager
    participant A1 as Attention N+1

    A0->>W: 设置当前 group
    A0->>A0: forward 并注册符合 policy 的 saved tensors
    A0->>W: forward hook 发起 D2H
    A1->>W: 等待 previous group 的 D2H
    A1->>A1: forward
    A1->>W: backward pre-hook 预取 Attention N
    A1->>A1: backward
    A0->>W: 等待 H2D 完成后 backward

7. 组件依赖

依赖组件 强依赖 / 弱依赖 当前状态 未 ready 时本期能力
PyTorch saved tensor hooks 强依赖 已有 无法提供本特性
swap_wrapper / SwapManager 强依赖 已有 无法提供本特性
FSDP2 弱依赖 使用现有能力 可先完成无 FSDP Level0;正式验收需覆盖目标组合
MindSpore 不涉及 本期不支持 明确报错
图编译 不涉及 本期不支持 明确报错

本期最小可交付能力:Torch + HF Qwen3-30B-A3B 在非 PP 两种场景下可通过同一配置开启 Attention swap。

8. 约束与兼容性

类型 内容
不支持项 MindSpore、图编译、非 Qwen3-30B-A3B 模型、PP 场景、重算
默认兼容性 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、group copy 降低影响,需输出实测数据
Transformers 版本 依赖 Qwen3MoeForCausalLM -> model.layers[*].self_attn 边界;结构不匹配时快速失败,不能静默运行

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 + 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

likedislike
DavidFFFanDavidFFFan成员
8月10日 将 silkage_jiajia 设为负责人
此处折叠了11条事件消息 查看更多
songjiaqisongjiaqi成员
8月11日 修改了issue 的描述
hui-zhang940
8月12日 评论:

功能接口配置只能通过模型yaml配置么?

likedislike
hui-zhang940
8月12日 评论:

功能接口配置只能通过模型yaml配置么?

@hui-zhang940

与hyper_parallel端到端训练流程一起串讲,由模型团队加上该特性进行统一验收

likedislike
songjiaqisongjiaqi成员
8月13日 修改了issue 的描述
此处折叠了5条事件消息 查看更多
songjiaqisongjiaqi成员
17 天前 修改了issue 的描述