已关闭
[RFC]: SwapLayerInput RFC #176
JialiZheng创建于  6月15日关闭于  7月11日
JialiZheng
JialiZheng成员
6月15日 创建

💻 需求背景、当前现状、期望实现的功能内容、具体的设计方案、以及测试方案

状态 (Status): Draft
作者 (Authors): @JialiZheng1
创建日期 (Created): 2026-06-15
更新日期 (Updated): 2026-06-15
相关 Issue/PR: https://gitcode.com/Ascend/MindSpeed/pull/3529


1. 概述

1.1 简介

本提案提出一种在 NPU 上训练大语言模型时,通过将 Transformer Layer 的输入张量(hidden_states)在层间执行 Device↔Host 交换(Swap),有效降低显存峰值占用的方案 —— SwapLayerInput。

核心思路是:在前向传播过程中,当前层计算完成后,将上一层的输入张量从 NPU 显存异步拷贝至 Host 内存并释放显存;在反向传播过程中,提前将所需的输入张量从 Host 预取回 NPU 显存。通过流水线式的异步 D2H/H2D 传输,将显存占用的时间窗口从"整个训练生命周期"缩短为"仅当前层需要时",从而在几乎不影响训练吞吐的前提下显著降低显存需求。

1.2 动机

在大模型训练场景中,Transformer Layer 数量可达数十甚至上百层。典型训练流程中,重计算可以将每层前向的中间激活释放,但是每一层的输入张量(hidden_states)还是需要保留到反向传播时计算梯度,导致显存占用与层数线性增长,成为制约模型规模扩展的主要瓶颈。

当前痛点:

  • 每层 hidden_states 需要驻留在 NPU 显存中直到反向传播完成,导致峰值显存占用极高;
  • 显存不足时只能减小 batch_size 或序列长度,影响训练效率;

本方案通过在层间异步交换输入张量,将显存中同时驻留的 hidden_states 数量从 O(n) 降为 O(1),以极小的传输开销换取大幅显存节省。

1.3 目标

目标:

  • 将 Transformer 训练中 hidden_states 的显存占用从 O(num_layers × batch_size × seq_len × hidden_dim) 降低为 O(batch_size × seq_len × hidden_dim);
  • 通过异步 NPU Stream 实现 D2H/H2D 传输与计算重叠,最小化吞吐损失;
  • 支持 Pipeline Parallel (PP) 场景下的多 micro-batch FIFO 管理;
  • 提供装饰器风格的 API,对已有 Transformer Layer 代码侵入性最小。

非目标:

  • 不涉及 Optimizer State、Gradient、Parameter 等其他张量的显存优化;
  • 不涉及 CPU 内存的高效管理策略(如 NUMA 绑定、内存池复用等),仅使用 pinned memory 保证传输效率;

2. 用例分析

场景 描述 关键要求
标准 Transformer 训练 前向传播中将 hidden_states 异步 D2H,反向传播前异步 H2D 预取 异步传输与计算重叠,吞吐损失 < 5%
Pipeline Parallel (PP) 多 micro-batch 按 FIFO 顺序管理 SwapTensors 队列 保证 PP 调度顺序正确,不同 micro-batch 的 SwapTensors 不混淆
FBOverlap (MoE) 前向/反向计算与通信 overlap 场景下的 Swap 调度 兼容 FBOverlap 的 checkpoint 与 1F1B 调度模式
MTP (Multi-Token Prediction) MTP 模块中同一 TransformerLayer 被多次调用 正确识别 MTP 调用,避免重复注册 SwapManager

使用限制:

  • 需要 Host 端有足够的 pinned memory(约为 hidden_states 总大小);

3. 方案设计

3.1 总体方案

3.1.1 架构概览

┌─────────────────────────────────────────────────────────────┐
│                  Transformer Layer Pipeline                   │
│                                                               │
│  Layer i-1          Layer i           Layer i+1              │
│  ┌─────────┐       ┌─────────┐       ┌─────────┐            │
│  │ Forward  │──HS──▶│ Forward  │──HS──▶│ Forward  │           │
│  │  (done)  │       │(running) │       │(pending) │           │
│  └────┬─────┘       └────┬─────┘       └─────────┘           │
│       │                  │                                    │
│       │ D2H (async)      │                                    │
│       ▼                  │                                    │
│  ┌─────────┐             │                                    │
│  │  Host   │◀────────────┘                                    │
│  │ Memory  │────▶ H2D (async, at backward)                    │
│  └─────────┘                                                  │
└─────────────────────────────────────────────────────────────┘

3.1.2 核心状态机

每个 SwapTensors 实例跟踪一组张量的传输状态:

  "device" ──swap_to_host()──▶ "d2h" ──wait_d2h()──▶ "host"
                                                     │
     ▲                                                │
     │         swap_to_device()                       │
     │                                                ▼
  "device" ◀──wait_h2d()── "h2d" ◀───────────────────┘

3.1.3 核心流程

前向传播(Forward):

  1. Layer i 的前向开始前,通过装饰器标记 hidden_states 的 swap_this_tensor = True
  2. 若在 torch.no_grad() 模式下,将 hidden_states 注册到当前层的 SwapLayerInputManager.batch_stack 中;
  3. forward_hook 在 Layer i 计算完成后触发:等待 Layer i-1 的 D2H 传输完成并释放 NPU 显存,同时启动当前层 hidden_states 的异步 D2H;
  4. 最后一层不执行 swap(不需要为下一层节省显存)。

反向传播(Backward):

  1. 通过 register_hook 注册反向梯度 Hook;
  2. backward_hook 触发时:等待当前层的 H2D 预取完成(若当前层不是最后一层),同时为前一层启动异步 H2D 预取;
  3. 完成 H2D 后释放 CPU 端 pinned memory。

3.1.4 Pipeline Parallel 支持

SwapLayerInputManager 内部维护一个 batch_stack: List[SwapTensors](FIFO 队列),每个 micro-batch 对应一个 SwapTensors 条目。在 PP 调度中,多个 micro-batch 会依次经过同一层,FIFO 队列确保 swap 操作的顺序与 micro-batch 调度顺序一致。

3.2 技术选型

方案 优势 劣势 选择
装饰器注入(本方案) 对现有代码侵入性小,解耦 Swap 逻辑与 Layer 计算逻辑 依赖张量属性标记 swap_this_tensor ✅ 采用
手动在每个 Layer 中嵌入 Swap 代码 控制粒度更细 代码侵入性强,维护成本高 ❌ 放弃
使用 PyTorch Hook 机制自动拦截 完全无侵入 Hook 执行时机可能难以精确控制,调试困难 ❌ 放弃

技术选型理由:装饰器模式在 MindSpeed 项目中已有成熟实践(如 FBOverlap),且能保证 Swap 逻辑与 Layer 实现的清晰边界。

3.3 功能与性能设计

3.3.1 模块结构

mindspeed/core/memory/swap_layer_input/
├── __init__.py                       # 空文件,标识为 Python 包
├── swap_layer_input_manager.py       # 核心实现:SwapTensors、SwapLayerInputManager
└── swap_layer_input.py               # 装饰器/Wrapper 层,提供给 Transformer Layer 使用

3.3.2 核心类与函数

is_valid_for_swap(tensor, custom_check_fn):校验张量是否可被交换。

  • 排除 nn.Parameter(模型参数不应被 swap);
  • 排除空存储张量;
  • 支持自定义校验函数。

SwapTensors:管理一组张量的 Device↔Host 传输。

  • 维护每个张量的 CPU pinned memory 副本 (tensor_cpus);
  • 追踪原始 storage size,支持 slice tensor(storage.size() ≠ numel());
  • 通过独立 NPU Stream 执行异步 D2H/H2D,与计算 Stream 重叠。

SwapLayerInputManager:全局交换管理器。

  • 通过 manager_map 类变量维护所有层的 Manager 实例,支持 module_tag 区分不同模块组;
  • 使用 layer_idx 确定当前层在 Pipeline 中的位置;
  • 维护 batch_stack FIFO 队列支持 PP 多 micro-batch;
  • 使用两条 NPU Stream(_d2h_stream_h2d_stream)分别处理两个方向的传输。

装饰器函数(swap_layer_input.py):

装饰器 用途
swap_layer_input_init_wrapper 在 Layer __init__ 中注册 SwapLayerInputManager
swap_layer_input_forward_wrapper 标准前向传播的 Swap 逻辑(标记张量、swap-out、注册 backward hook)
swap_layer_input_fboverlap_forward_wrapper FBOverlap 场景的前向 Swap(支持 checkpoint 分支)
swap_layer_input_fboverlap_1f1b_wrapper FBOverlap 1F1B 场景的 Swap(同时管理前向层和后向层的 swap)
swap_layer_input_fboverlap_backward_wrapper FBOverlap 场景的反向 Swap

3.3.3 性能设计

  • 异步传输:D2H 和 H2D 使用独立 NPU Stream,与计算 Stream 并行,理想情况下传输完全被计算覆盖;
  • 批量操作:以 SwapTensors 组为单位批量执行 storage.copy_,减少 kernel launch 开销;
  • Pinned Memory:CPU 端使用 pin_memory=True 分配,保证 PCIe/NPU 互联链路的 DMA 传输效率;
  • 懒初始化:Stream 在首次使用时创建(_ensure_streams),避免不必要的资源占用。

3.4 安全隐私与 DFX 设计

兼容性

  • 兼容 MTP(Multi-Token Prediction)中同一 Layer 被多次调用的场景,通过 is_mtp 属性识别;
  • 向后兼容:未注册 swap_manager 的 Layer 直接执行原始前向逻辑,不影响功能。

可维护性

  • 装饰器模式将 Swap 逻辑与 Layer 计算逻辑解耦,便于独立修改和测试;
  • SwapTensors 状态机语义清晰(device → d2h → host → h2d → device),状态转换有明确的前置条件检查;
  • 模块化设计,manager 层负责状态管理与传输调度,wrapper 层负责与 Layer 的集成。

可测试性

  • SwapTensorsSwapLayerInputManager 可独立于 Transformer Layer 进行单元测试;
  • 状态检查(stat 属性)可在测试中验证状态转换的正确性;
  • custom_check_fn 参数允许测试注入自定义张量校验逻辑。

可靠性

  • 每次状态转换前检查当前状态,避免重复操作(如 swap_to_host 仅在 stat == "device" 时执行);
  • 使用 torch.npu.Event 确保 Stream 间的正确同步;

3.5 编程与调用设计

3.5.1 编程模型基本设计

开发环境设计:

  • 软件:PyTorch + torch_npu
  • 编程语言:Python 3.x

开发约束:

  • 张量交换前需要标记 swap_this_tensor = True 属性;
  • 装饰器必须按正确顺序应用(__init__forward)。

可验收设计:

  • 功能验收:在标准 Transformer 训练中,启用 Swap 后显存峰值降低比例 ≥ (num_layers - 1) / num_layers;
  • 性能验收:启用 Swap 后训练吞吐下降 ≤ 5%。

3.5.2 接口定义与设计

3.5.2.1 swap_layer_input_init_wrapper

  • 接口描述: 装饰 Transformer Layer 的 __init__ 方法,为每个 Layer 实例创建并注册 SwapLayerInputManager

  • 接口原型:

    def swap_layer_input_init_wrapper(fn: Callable) -> Callable
    
  • 输入/输出参数:

    参数名称 输入/输出 类型 描述 取值范围
    fn 输入 Callable 被装饰的 __init__ 函数 任意可调用对象
  • 返回参数:

    参数名称 类型 描述 取值范围
    wrapper Callable 包装后的 __init__ 函数
  • 异常处理: 无显式异常处理,异常由原始 __init__ 函数抛出。

  • 约束说明: 需要 Layer 类不包含 is_mtp 属性或 is_mtp=False,MTP 场景下会额外处理。

3.5.2.2 swap_layer_input_forward_wrapper

  • 接口描述: 装饰标准 Transformer Layer 的 forward 方法,实现 hidden_states 的标记、异步 swap-out 及 backward hook 注册。

  • 接口原型:

    def swap_layer_input_forward_wrapper(fn: Callable) -> Callable
    
  • 输入/输出参数:

    参数名称 输入/输出 类型 描述 取值范围
    fn 输入 Callable 被装饰的 forward 函数 任意可调用对象
  • 返回参数:

    参数名称 类型 描述 取值范围
    wrapper Callable 包装后的 forward 函数
  • 异常处理: 若 Layer 未注册 swap_manager,直接执行原始 forward 逻辑。

  • 约束说明: hidden_states 需为 torch.Tensor 类型,且位于 kwargs 的 hidden_states 键或 args[0] 位置。

3.5.2.3 swap_layer_input_fboverlap_1f1b_wrapper

  • 接口描述: 装饰 FBOverlap 1F1B 调度模式下的 forward 方法,同时管理前向层和后向层的 swap 操作。
  • 接口原型:
    def swap_layer_input_fboverlap_1f1b_wrapper(fn: Callable) -> Callable
    
  • 约束说明: 需要 kwargs 中包含 bwd_layer_graph(或 args[3] 为 LayerGraph 类型),用于获取反向层的 swap_manager。

3.5.2.4 swap_layer_input_fboverlap_backward_wrapper

  • 接口描述: 装饰 FBOverlap 场景下的 backward 函数。
  • 接口原型:
    def swap_layer_input_fboverlap_backward_wrapper(fn: Callable) -> Callable
    
  • 约束说明: 通过 args[1].layer 获取当前层的 swap_manager 实例。

3.5.3 编程手册设计

建议在 MindSpeed 已有文档中新增 "SwapLayerInput 使用指南" 章节,包含以下内容:

  1. 概述:SwapLayerInput 是什么、解决什么问题;
  2. 快速开始:在 Transformer Layer 中应用装饰器的最小示例;
  3. 装饰器说明:各 wrapper 的适用场景、前提条件和调用方式;
  4. PP/FBOverlap/MTP 场景的特别说明;
  5. 性能调优建议;
  6. 常见问题(FAQ)。

4. 缺点和风险

风险 影响 应对措施
Host 内存不足 pinned memory 分配失败导致训练中断 文档中说明 Host 内存需求(= num_layers × hidden_states_size),建议部署前校验
D2H/H2D 传输延迟 若传输时间大于单层计算时间,则无法完全 overlap,吞吐下降 通过 Stream 异步 + pinned memory 最大化带宽利用;对于小 hidden_size 场景,可通过 prefetch 参数调节
张量 slice 场景兼容 storage().size() ≠ numel() 的张量需特殊处理 SwapTensors 通过 is_slice_tensors 标记分别处理 slice 和非 slice 张量
MTP 场景 MTP 模块中同一 Layer 被调用两次,可能重复注册 Manager 通过 is_mtp 属性识别并移除多余注册,见 swap_layer_input_init_wrapper
FBOverlap 兼容性 多种调度模式(standard/checkpoint/1F1B)需不同 swap 策略 提供 5 个不同 wrapper 适配不同场景

5. 现有技术

  • DeepSpeed ZeRO-Offload:将 Optimizer State 和 Gradient offload 到 CPU,本方案关注的是 Activation(hidden_states)的 offload,两者互补;

6. 未解决问题

暂无


附录

  • 参考资料链接。
  • 术语表:
    • D2H:Device to Host,从 NPU 显存拷贝到 CPU 内存
    • H2D:Host to Device,从 CPU 内存拷贝到 NPU 显存
    • PP:Pipeline Parallel,流水线并行
    • MTP:Multi-Token Prediction,多 Token 预测
    • FBOverlap:Forward-Backward Overlap,前向-反向计算重叠优化
  • 文档更新计划: 在 MindSpeed 开发文档中新增 SwapLayerInput 使用指南章节。

替代方案

补充说明

欢迎加入社区,感谢您对社区的贡献 🎉!

likedislike
JialiZhengJialiZheng成员
6月15日 添加了label:rfc
JialiZhengJialiZheng成员
6月15日 修改了issue 的描述
Ggp513成员
6月22日 issue类型由 Bug-Report 改变为 RFC
ascend-robotascend-robot成员
7月3日 关联了看板:MindStudio ISSUE管理
Ggp513成员
7月11日 issue状态由 TODO 改变为 DONE
Ggp513成员
7月11日 关闭了 issue
ascend-robotascend-robot成员
7月11日 添加了label:resolved