状态 (Status): Draft 作者 (Authors): @JialiZheng1 创建日期 (Created): 2026-06-15 更新日期 (Updated): 2026-06-15 相关 Issue/PR: https://gitcode.com/Ascend/MindSpeed/pull/3529
本提案提出一种在 NPU 上训练大语言模型时,通过将 Transformer Layer 的输入张量(hidden_states)在层间执行 Device↔Host 交换(Swap),有效降低显存峰值占用的方案 —— SwapLayerInput。
核心思路是:在前向传播过程中,当前层计算完成后,将上一层的输入张量从 NPU 显存异步拷贝至 Host 内存并释放显存;在反向传播过程中,提前将所需的输入张量从 Host 预取回 NPU 显存。通过流水线式的异步 D2H/H2D 传输,将显存占用的时间窗口从"整个训练生命周期"缩短为"仅当前层需要时",从而在几乎不影响训练吞吐的前提下显著降低显存需求。
在大模型训练场景中,Transformer Layer 数量可达数十甚至上百层。典型训练流程中,重计算可以将每层前向的中间激活释放,但是每一层的输入张量(hidden_states)还是需要保留到反向传播时计算梯度,导致显存占用与层数线性增长,成为制约模型规模扩展的主要瓶颈。
当前痛点:
本方案通过在层间异步交换输入张量,将显存中同时驻留的 hidden_states 数量从 O(n) 降为 O(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) │ │ └─────────┘ │ └─────────────────────────────────────────────────────────────┘
每个 SwapTensors 实例跟踪一组张量的传输状态:
SwapTensors
"device" ──swap_to_host()──▶ "d2h" ──wait_d2h()──▶ "host" │ ▲ │ │ swap_to_device() │ │ ▼ "device" ◀──wait_h2d()── "h2d" ◀───────────────────┘
前向传播(Forward):
swap_this_tensor = True
torch.no_grad()
SwapLayerInputManager.batch_stack
forward_hook
反向传播(Backward):
register_hook
backward_hook
SwapLayerInputManager 内部维护一个 batch_stack: List[SwapTensors](FIFO 队列),每个 micro-batch 对应一个 SwapTensors 条目。在 PP 调度中,多个 micro-batch 会依次经过同一层,FIFO 队列确保 swap 操作的顺序与 micro-batch 调度顺序一致。
SwapLayerInputManager
batch_stack: List[SwapTensors]
swap_this_tensor
技术选型理由:装饰器模式在 MindSpeed 项目中已有成熟实践(如 FBOverlap),且能保证 Swap 逻辑与 Layer 实现的清晰边界。
mindspeed/core/memory/swap_layer_input/ ├── __init__.py # 空文件,标识为 Python 包 ├── swap_layer_input_manager.py # 核心实现:SwapTensors、SwapLayerInputManager └── swap_layer_input.py # 装饰器/Wrapper 层,提供给 Transformer Layer 使用
is_valid_for_swap(tensor, custom_check_fn):校验张量是否可被交换。
is_valid_for_swap(tensor, custom_check_fn)
nn.Parameter
SwapTensors:管理一组张量的 Device↔Host 传输。
tensor_cpus
SwapLayerInputManager:全局交换管理器。
manager_map
module_tag
layer_idx
batch_stack
_d2h_stream
_h2d_stream
装饰器函数(swap_layer_input.py):
swap_layer_input.py
swap_layer_input_init_wrapper
__init__
swap_layer_input_forward_wrapper
swap_layer_input_fboverlap_forward_wrapper
swap_layer_input_fboverlap_1f1b_wrapper
swap_layer_input_fboverlap_backward_wrapper
pin_memory=True
_ensure_streams
is_mtp
swap_manager
manager
wrapper
stat
custom_check_fn
swap_to_host
stat == "device"
torch.npu.Event
开发环境设计:
开发约束:
forward
可验收设计:
接口描述: 装饰 Transformer Layer 的 __init__ 方法,为每个 Layer 实例创建并注册 SwapLayerInputManager。
接口原型:
def swap_layer_input_init_wrapper(fn: Callable) -> Callable
输入/输出参数:
返回参数:
异常处理: 无显式异常处理,异常由原始 __init__ 函数抛出。
约束说明: 需要 Layer 类不包含 is_mtp 属性或 is_mtp=False,MTP 场景下会额外处理。
is_mtp=False
接口描述: 装饰标准 Transformer Layer 的 forward 方法,实现 hidden_states 的标记、异步 swap-out 及 backward hook 注册。
def swap_layer_input_forward_wrapper(fn: Callable) -> Callable
异常处理: 若 Layer 未注册 swap_manager,直接执行原始 forward 逻辑。
约束说明: hidden_states 需为 torch.Tensor 类型,且位于 kwargs 的 hidden_states 键或 args[0] 位置。
torch.Tensor
hidden_states
def swap_layer_input_fboverlap_1f1b_wrapper(fn: Callable) -> Callable
kwargs
bwd_layer_graph
LayerGraph
def swap_layer_input_fboverlap_backward_wrapper(fn: Callable) -> Callable
args[1].layer
建议在 MindSpeed 已有文档中新增 "SwapLayerInput 使用指南" 章节,包含以下内容:
prefetch
is_slice_tensors
暂无
附录
无
欢迎加入社区,感谢您对社区的贡献 🎉!
💻 需求背景、当前现状、期望实现的功能内容、具体的设计方案、以及测试方案
状态 (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 数量从 O(n) 降为 O(1),以极小的传输开销换取大幅显存节省。
1.3 目标
目标:
非目标:
2. 用例分析
使用限制:
3. 方案设计
3.1 总体方案
3.1.1 架构概览
3.1.2 核心状态机
每个
SwapTensors实例跟踪一组张量的传输状态:3.1.3 核心流程
前向传播(Forward):
swap_this_tensor = True;torch.no_grad()模式下,将 hidden_states 注册到当前层的SwapLayerInputManager.batch_stack中;forward_hook在 Layer i 计算完成后触发:等待 Layer i-1 的 D2H 传输完成并释放 NPU 显存,同时启动当前层 hidden_states 的异步 D2H;反向传播(Backward):
register_hook注册反向梯度 Hook;backward_hook触发时:等待当前层的 H2D 预取完成(若当前层不是最后一层),同时为前一层启动异步 H2D 预取;3.1.4 Pipeline Parallel 支持
SwapLayerInputManager内部维护一个batch_stack: List[SwapTensors](FIFO 队列),每个 micro-batch 对应一个SwapTensors条目。在 PP 调度中,多个 micro-batch 会依次经过同一层,FIFO 队列确保 swap 操作的顺序与 micro-batch 调度顺序一致。3.2 技术选型
swap_this_tensor技术选型理由:装饰器模式在 MindSpeed 项目中已有成熟实践(如 FBOverlap),且能保证 Swap 逻辑与 Layer 实现的清晰边界。
3.3 功能与性能设计
3.3.1 模块结构
3.3.2 核心类与函数
is_valid_for_swap(tensor, custom_check_fn):校验张量是否可被交换。nn.Parameter(模型参数不应被 swap);SwapTensors:管理一组张量的 Device↔Host 传输。tensor_cpus);SwapLayerInputManager:全局交换管理器。manager_map类变量维护所有层的 Manager 实例,支持module_tag区分不同模块组;layer_idx确定当前层在 Pipeline 中的位置;batch_stackFIFO 队列支持 PP 多 micro-batch;_d2h_stream、_h2d_stream)分别处理两个方向的传输。装饰器函数(
swap_layer_input.py):swap_layer_input_init_wrapper__init__中注册SwapLayerInputManagerswap_layer_input_forward_wrapperswap_layer_input_fboverlap_forward_wrapperswap_layer_input_fboverlap_1f1b_wrapperswap_layer_input_fboverlap_backward_wrapper3.3.3 性能设计
SwapTensors组为单位批量执行 storage.copy_,减少 kernel launch 开销;pin_memory=True分配,保证 PCIe/NPU 互联链路的 DMA 传输效率;_ensure_streams),避免不必要的资源占用。3.4 安全隐私与 DFX 设计
兼容性
is_mtp属性识别;swap_manager的 Layer 直接执行原始前向逻辑,不影响功能。可维护性
SwapTensors状态机语义清晰(device → d2h → host → h2d → device),状态转换有明确的前置条件检查;manager层负责状态管理与传输调度,wrapper层负责与 Layer 的集成。可测试性
SwapTensors和SwapLayerInputManager可独立于 Transformer Layer 进行单元测试;stat属性)可在测试中验证状态转换的正确性;custom_check_fn参数允许测试注入自定义张量校验逻辑。可靠性
swap_to_host仅在stat == "device"时执行);torch.npu.Event确保 Stream 间的正确同步;3.5 编程与调用设计
3.5.1 编程模型基本设计
开发环境设计:
开发约束:
swap_this_tensor = True属性;__init__→forward)。可验收设计:
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输入/输出参数:
__init__函数返回参数:
__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输入/输出参数:
forward函数返回参数:
forward函数异常处理: 若 Layer 未注册
swap_manager,直接执行原始 forward 逻辑。约束说明: hidden_states 需为
torch.Tensor类型,且位于 kwargs 的hidden_states键或 args[0] 位置。3.5.2.3 swap_layer_input_fboverlap_1f1b_wrapper
def swap_layer_input_fboverlap_1f1b_wrapper(fn: Callable) -> Callablekwargs中包含bwd_layer_graph(或 args[3] 为LayerGraph类型),用于获取反向层的 swap_manager。3.5.2.4 swap_layer_input_fboverlap_backward_wrapper
def swap_layer_input_fboverlap_backward_wrapper(fn: Callable) -> Callableargs[1].layer获取当前层的 swap_manager 实例。3.5.3 编程手册设计
建议在 MindSpeed 已有文档中新增 "SwapLayerInput 使用指南" 章节,包含以下内容:
4. 缺点和风险
prefetch参数调节SwapTensors通过is_slice_tensors标记分别处理 slice 和非 slice 张量is_mtp属性识别并移除多余注册,见swap_layer_input_init_wrapper5. 现有技术
6. 未解决问题
暂无
附录
替代方案
无
补充说明
无
欢迎加入社区,感谢您对社区的贡献 🎉!