已开启
[RFC] Torch Qwen3-30B-A3B Attention 激活内存 Swap #323
DavidFFFan创建于 8月10日
8月10日 将 silkage_jiajia 设为负责人
此处折叠了11条事件消息 查看更多
8月11日 修改了issue 的描述
hui-zhang940
8月12日 评论:
8月12日 评论:
功能接口配置只能通过模型yaml配置么?


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


8月13日 修改了issue 的描述
此处折叠了5条事件消息 查看更多
17 天前 修改了issue 的描述
1. 基本信息
distributed/attention_swap.py、models/_transformers/model_builder.py、core/activation_checkpoint、trainer2. 背景
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 目标
Qwen3MoeForCausalLM,验收模型为 Qwen3-30B-A3B。Qwen3MoeAttention范围内 autograd 保存用于 backward 的激活,不 swap DecoderLayer 的两个 RMSNorm、residual 和 MLP/MoE 激活。SwapManager.set_forward_prefetch_layer()按 Attention 扫描顺序建立非 PP 的异步D2H/H2D 调度。
swap_wrapper。activation_swap=none。3.2 非目标
torch.compile或其他图编译;swap attention 与图编译同时开启时应明确报错。attn、attention或其他名称的 Attention 模块;模型没有self_attn子模块时会快速失败。
4. 相关实现参考
swap_wrapper/saved_tensors_hooksSwapManager.set_forward_prefetch_layerQwen3MoeDecoderLayer.self_attn边界清晰self_attn,避免重写 forward,降低版本耦合5. 对外接口
5.1 接口定义
在新 Trainer 的公共配置中增加一个开关:
activation_swap: Literal["none", "attention"] = "none"YAML 示例:
activation_swap: attentionactivation_swapstrnonenone/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=noneattention+pp_size=1+ checkpoint off + compile offattention+ full/selective checkpointValueErrorattention+torch.compileValueErrorself_attnself_attnself_attnValueError6. 方案设计
6.1 Attention 目标发现
从根模型开始递归调用
named_children():self_attn时,将其记录为目标,并停止继续递归该目标内部;id()对共享模块去重;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 保留在设备侧:
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 完成后 backward7. 组件依赖
swap_wrapper/SwapManager本期最小可交付能力:Torch + HF Qwen3-30B-A3B 在非 PP 两种场景下可通过同一配置开启 Attention swap。
8. 约束与兼容性
activation_swap=none时不包装、不注册 hook,行为与当前版本一致(stage, microbatch)调度(目前不支持);非 PP 按 Attention 层调度Qwen3MoeForCausalLM -> model.layers[*].self_attn边界;结构不匹配时快速失败,不能静默运行9. 验证设计
9.1 用例分层
9.2 核心正确性验证
torch.testing.assert_close;BF16 设备测试沿用项目既有精度阈值。9.3 交互验证
9.4 性能和显存验证
activation_swap=noneattention性能不预设未经实测的固定收益比例。PR 验收材料必须同时给出:模型配置、序列长度、global/micro batch、并行策略、rank 数、Attention 实现、基线/开启后的显存和吞吐。
Qwen3-30B-A3B seq_length:1024
noneattention10. 实现计划与工作量
swap_wrapperSwapManager和 PP schedule总计:使用 AI 辅助约 1.25~2.5 人天。其中核心代码和 UT 约 0.75~1.25 人天,真实设备验证约 0.5~1.25 人天。设备排队和环境问题不计入纯开发工时。
11. 验收 Checklist