swap
checkpoint
optimizer
目前的训练中Adam/AdamW优化器状态会长期驻留在设备显存中,对大模型而言,显存容量成为更早、更明显的瓶颈。 Swap optimizer 的核心思路是:在两次 optimizer update 之间,将 Adam/AdamW 的大状态张量保存在 pinned CPU 内存中;只在执行 update 时把当前批次状态搬到设备,借助独立拷贝流尽量把下一批 H2D、当前批计算和上一批 D2H 重叠,从而降低 optimizer step 期间的设备峰值内存。
1. 支持优化器范围:torch: 原生 Adam、AdamW、HP自研 AdamW;mindspore: 原生Adam、AdamWeightDecay、MF自研 AdamW; 2. 支持swap optimizer特性与FSDP,TP/PP/EP/CP,重计算,swap activation等特性叠加使用。
1. 对已支持优化器的使用有限制,以便以batch为粒度更新; 2. 本期不支持 muon 优化器,待后期补全。
包装已有 Adam/AdamW optimizer,返回当前框架对应的 swap optimizer。之后继续按原框架方式使用:
optimizer.step()
optimizer.zero_grad()
optimizer(gradients)
约束:
torch.optim.Adam
torch.optim.AdamW
AdamW
nn.Adam
nn.AdamWeightDecay
config=None
packed_swap=True
packed_swap=False
swap_times
16
> 0
state_keys
None
exp_avg
exp_avg_sq
max_exp_avg_sq
master_param
min_numel
1024
>= 0
include_master_params
False
packed_swap
True
其他说明:
state_keys=None
packed_swap=True/False
from hyper_parallel.core.optimizer import SwapOptimizerConfig, swap_optimizer optimizer = swap_optimizer( base_optimizer, SwapOptimizerConfig( swap_times=16, min_numel=1024, packed_swap=True, ), )
class torch.optim.Adam(params, lr=0.001, betas=(0.9, 0.999), eps=1e-08, weight_decay=0, amsgrad=False, *, foreach=None, maximize=False, capturable=False, differentiable=False, fused=None, decoupled_weight_decay=False) class torch.optim.AdamW(params, lr=0.001, betas=(0.9, 0.999), eps=1e-08, weight_decay=0.01, amsgrad=False, *, maximize=False, foreach=None, capturable=False, differentiable=False, fused=None) Optimizer.step(closure: None = None) → None[source]
foreach=False
fused=False
capturable=False
differentiable=False
swap_optimizer
_fused_adam_
step()
closure
hyper_parallel.core.optimizer.adamw( params: List[torch.Tensor], grads: List[torch.Tensor], exp_avgs: List[torch.Tensor], exp_avg_sqs: List[torch.Tensor], max_exp_avg_sqs: List[torch.Tensor], step: int, *, amsgrad: bool, beta1: float, beta2: float, lr: float, weight_decay: float, eps: float, maximize: bool ) Optimizer.step(closure: None = None)
class mindspore.nn.Adam(params, learning_rate=1e-3, beta1=0.9, beta2=0.999, eps=1e-8, use_locking=False, use_nesterov=False, weight_decay=0.0, loss_scale=1.0, use_amsgrad=False, **kwargs) class mindspore.nn.AdamWeightDecay(params, learning_rate=1e-3, beta1=0.9, beta2=0.999, eps=1e-6, weight_decay=0.0)
use_lazy=False
use_offload=False
class AdamW( params, learning_rate=1e-3, betas=(0.9, 0.999), eps=1e-8, weight_decay=0.0, enable_cpu_offload=False, enable_fused_opt=False, use_fused=False, **kwargs )
enable_cpu_offload=False
以每个 optimizer state 为独立单位执行搬运。 优点:适配范围广、实现灵活。 缺点:小块拷贝和显存分配次数较多,调度开销较大。 per tensor swap 掩盖关系:
将相同 dtype 的多个状态 tensor 打包进连续的 pinned CPU buffer,每step使用两个可复用的 device staging buffer。 优点:把大量小拷贝合并成连续拷贝,减少 allocation 和调度开销,并通过 A/B 双缓冲重叠状态传输与参数更新。 缺点:使用有限制 packed swap 掩盖关系:
torch 侧 packed swap 内存结构:
mindspore 侧 packed swap 内存结构:
sequenceDiagram autonumber participant Train as Training Loop participant Wrapper as TorchSwapOptimizer participant Adapter as TorchAdamBaseAdapter participant Runtime as PipelineSwapRuntime participant Copy as Copy Stream participant CPU as CPU Pinned Memory participant Compute as Compute Stream participant Param as Model Parameter Note over Wrapper,CPU: 初始化阶段,config.packed_swap=False Train->>Wrapper: swap_optimizer(base_optimizer, config) Wrapper->>Adapter: 创建 optimizer adapter Adapter->>Runtime: 遍历 param_groups(处理目前已存在state,如checkpoint加载),<br>注册 exp_avg / exp_avg_sq 等 SwapSlot Wrapper->>Runtime: offload_initial_slots(initial_slots) Runtime->>CPU: 为所有已存在的 optimizer state 创建/获取 pinned CPU mirror Runtime->>CPU: D2H copy optimizer state Runtime->>Runtime: 释放 device state storage Note over Train,Param: 一次 optimizer.step() Train->>Wrapper: optimizer.step() Wrapper->>Adapter: prepare_step() loop 遍历有梯度的 parameter Adapter->>Adapter: 懒初始化 optimizer state: <br>创建 pinned CPU mirror -> 创建 device tensor 以构建 SwapSlot 后随即 resize(0) Adapter->>Adapter: 为当前参数的每个 state 找到/创建 SwapSlot Adapter->>Adapter: 获取 slots 创建 UpdateUnit(param, grad, slots) <br>(UpdateUnit 包含一个 param 的所有要做swap 的 optimizer states) end Adapter-->>Wrapper: UpdateUnit 列表 Wrapper->>Runtime: partition(units) Runtime->>Runtime: 按可 swap 的优化器状态的 state_nbytes 近似均衡切分 batch Wrapper->>Runtime: run_pipeline Note over Runtime,CPU: 首批预取 Runtime->>Copy: prefetch(batch 0) Runtime->>Runtime: 恢复 batch 0 每个 state tensor 的 device storage CPU->>Copy: H2D param 优化器状态(exp_avg、exp_avg_sq...) Copy->>Copy: 记录 batch 0 H2D ready event loop batch i Runtime->>Compute: wait batch i H2D complete event Runtime->>Runtime: wait (batch i-1) D2H complete event,释放 device state storage Runtime->>Copy: prefetch(batch i+1) <br> 恢复 device storage,做 H2D Runtime->>Adapter: step_batch(batch i) Adapter->>Compute: functional Adam/AdamW Compute->>Param: 原地更新 model parameter Compute->>Compute: 原地更新 exp_avg / exp_avg_sq Compute->>Copy: 记录 update-complete event Runtime->>Copy: offload(batch i):wait update-complete event, 做 D2H Copy->>Copy: 记录 D2H complete event end Runtime->>Runtime: wait last D2H event,释放 device state storage Runtime-->>Wrapper: pipeline 完成 Wrapper-->>Train: step() 返回 Note over Param,CPU: step 结束状态 Note over Param: Model parameter 已在设备上原地更新 Note over CPU: swappable moments 保存在 CPU pinned memory
sequenceDiagram autonumber participant Train as Training Loop participant Wrapper as TorchSwapOptimizer participant Adapter as Adam Adapter participant Runtime as Packed Runtime participant CPU as CPU Pinned Buffer participant Copy as Copy Stream participant A as Device Arena A participant B as Device Arena B participant Compute as Compute Stream participant Param as Model Parameter Note over Wrapper,CPU: 初始化阶段,config.packed_swap=True Train->>Wrapper: swap_optimizer(base_optimizer, config) Wrapper->>Adapter: 创建 optimizer adapter Adapter->>Runtime: 遍历 param_groups(处理目前已存在state,如checkpoint加载),<br>注册 exp_avg / exp_avg_sq 等 SwapSlot Wrapper->>Runtime: ofload_initial_slots(initial_slots) Runtime->>CPU: 为所有已存在的 optimizer state 创建/获取 pinned CPU mirror Runtime->>CPU: D2H copy optimizer state Runtime->>Runtime: 释放 device state storage Wrapper->>Runtime: prepare_packed_host(initial_slots) loop 遍历 dtype Runtime->>CPU: 按 dtype 创建连续 pinned CPU buffer Runtime->>CPU: 将 cpu_tensor 复制到连续 pinned CPU buffer 中对应的一段view Runtime->>Runtime: 注册每个 slot 的 host_offset,<br>设置 slot 的 cpu_tensor 为对应 host_view end Note over Train,Param: optimizer.step() Train->>Wrapper: step() Wrapper->>Adapter: prepare_step() loop 遍历有梯度的 parameter Adapter->>Adapter: 懒初始化 optimizer state: 创建 SwapSlot Adapter->>Adapter: 登记现有的 SwapSlot end Adapter->>Runtime: prepare_packed_host() alt packed_slots 完全没变 Runtime->>Runtime: 直接复用旧的连续 CPU buffer else packed_slots 有变化 loop Runtime->>CPU: 按 dtype 创建连续 pinned CPU buffer Runtime->>CPU: 将 cpu_tensor 或 tensor 复制到连续 pinned CPU buffer 中对应的一段view Runtime->>Runtime: 注册每个 slot 的 host_offset,<br>设置 slot 的 cpu_tensor 为对应 host_view end end Adapter->>Adapter: 获取 slots 创建 UpdateUnit(...) (UpdateUnit 包含一个 param 的所有要做swap 的 optimizer states) Adapter-->>Wrapper: 返回 UpdateUnit[] Wrapper->>Runtime: partition(units) Runtime->>Runtime: 按可 swap 的优化器状态的 state_nbytes 均衡切分 batch Wrapper->>Runtime: run_pipeline -> _run_packed_pipeline par begin_packed_step(batches) Runtime->>Runtime: 遍历 batches ,收集每 batch 需要 swap 的 packed slots,每 batch slots 按 dtype 分组 Runtime->>Runtime: 然后 batch 中 slots 根据 dtype 集中排列,组合成连续 region -> 为每个 batch 创建 PackedBatchPlan 记录 region 信息 Runtime->>Runtime: 统计在一 batch 中每种 dtype 所需最大的空间,<br>并用各 dtype 最大值计算一个布局,布局内各 dtype 区域按512对齐 Runtime->>A: 创建 staging arena A,按 dtype 创建 device views Runtime->>B: 创建 staging arena B,按 dtype 创建 device views end Note over Runtime,A: 准备 batch 0 Runtime->>Copy: prefetch(batch 0, arena A) Copy->>A: Batch 0 连续 region H2D Copy-->>Runtime: ready event Batch 0 Note over Runtime,B: 预取 batch 1 Runtime->>Copy: prefetch(batch 1, arena B) Copy->>B: Batch 1 连续 region H2D Copy-->>Runtime: ready event Batch1 loop 例如 batch i 使用 arena A Runtime->>Compute: wait batch i alt i >= 2 Runtime->>Compute: wait batch i-2 offload event Runtime->>CPU: 将 batch i-2 slots 重新绑定CPU view end Runtime->>A: 将 Batch i slots 绑定到 arena i%2 Runtime->>Adapter: step_batch(batch i) Adapter->>Compute: functional Adam/AdamW Compute->>Param: 原地更新模型参数 Compute->>Compute: 更新 staging 中的 exp_avg/exp_avg_sq... Compute-->>Copy: 记录 update-complete event Copy->>A: 等待 update-complete A->>CPU: batch i D2H CPU->>A: 同一 copy stream 上执行 batch i+2 H2D Copy-->>Runtime: offload&next-ready event end Note over Runtime,B: 最后2个 batch Copy->>Runtime: 等待最后的 offload event Runtime->>CPU: 将 slots 重新绑定 CPU view Note over Runtime,B: step mo末尾 Runtime->>A: 释放 arena A storage Runtime->>B: 释放 arena B storage Wrapper-->>Train: step() 返回
sequenceDiagram autonumber participant Train as Training Loop participant Wrapper as MindSporeSwapOptimizer participant Adapter as MindSpore Adam Adapter participant BaseOpt as Base Optimizer participant Runtime as MindSporeSwapRuntime participant Copy as Copy Stream participant CPU as CPU Pinned Memory participant Compute as Compute Stream participant Param as Model / Master Parameter Note over Wrapper,Adapter: MindSpore optimizer state 通常在 optimizer 构造时已经存在 Train->>Wrapper: swap_optimizer(base_optimizer, config) Wrapper->>Adapter: 创建 adapter Adapter->>BaseOpt: 遍历已创建的 moment1/moment2/vhat<br>或 exp_avg/exp_avg_sq/fp32_params 优化器状态,注册SwapSlot <br>(ms 的 optimizer 在构造时就把 optimizer state 创建好了) Wrapper->>Runtime: offload_initial_slots(initial_slots) Runtime->>CPU: 为所有已存在的 optimizer state 创建/获取 pinned CPU mirror Runtime->>CPU: D2H copy optimizer state Runtime->>Runtime: 释放 device state storage Note over Train,Param: 一次 optimizer update Train->>Wrapper: construct(gradients) Wrapper->>Adapter: prepare_step(gradients) Adapter->>BaseOpt: 校验/预处理梯度,获取 lr、weight decay<br>更新 global_step / beta powers loop 遍历有梯度的 parameter Adapter->>Adapter: 获取 state slots(初始化时已创建) Adapter->>Adapter: 创建 UpdateUnit(param, grad, selected slots) end Adapter-->>Wrapper: UpdateUnit 列表 Wrapper->>Runtime: partition(units) Runtime->>Runtime: 按 swappable slot.storage_nbytes<br>保持原顺序近似均衡切分 Wrapper->>Runtime: run_pipeline() Note over Runtime,CPU: 首批预取 Runtime->>Copy: prefetch(batch 0) Runtime->>Runtime: 恢复 batch 0 每个 state tensor 的 device storage CPU->>Copy: H2D param 优化器状态 Copy->>Copy: 记录 batch 0 H2D ready event loop batch i Runtime->>Compute: wait batch i H2D complete event Runtime->>Runtime: wait (batch i-1) D2H complete event,释放 device state storage Runtime->>Copy: prefetch(batch i+1) <br> 恢复 device storage,做 H2D Runtime->>Adapter: step_batch(batch i) alt mindspore.nn.Adam Adapter->>Compute: 每个 unit 调用 _apply_adam / AMSGrad opt Compute->>Param: 原地更新 model param 和 moments else mindspore.nn.AdamWeightDecay Adapter->>Compute: 每个 unit 调用 fused_opt Compute->>Param: 原地更新 model param 和 moments else MindFormers AdamW Adapter->>Compute: _run_adamw_opt / _run_fused_adamw_opt Compute->>Param: 原地更新 model param 和 moments opt include_master_params=True Adapter->>Param: 当前 batch fp32 master param 同步到 model param <br>(因为马上就会offload fp32 master param) end end Runtime->>Compute: record update-complete event Runtime->>Runtime: refresh_swappable_slots<br>(ms 有些优化器状态在第一次 step 后才会建立有效的 device storage) Runtime->>Copy: offload(batch i):wait update-complete event,做D2H Copy->>Copy: record D2H-complete event end Runtime->>Runtime: wait last D2H event,释放 device state storage Wrapper->>Adapter: finish_step() opt MindFormers 且 include_master_params=False Adapter->>Param: 将全部 fp32 main params 同步到 model params end Wrapper-->>Train: tuple(batch results) Note over Param,CPU: 更新结束 Note over CPU: swappable moments 由 pinned CPU mirror 持有 Note over Param: model param 已更新
sequenceDiagram autonumber participant Train as Training Loop participant Wrapper as MindSporeSwapOptimizer participant Adapter as MindFormers AdamW Adapter participant Runtime as Packed Runtime participant CPU as CPU Pinned Buffer participant Copy as Copy Stream participant A as Device Staging Arena A participant B as Device Staging Arena B participant Compute as Compute Stream participant Param as Model Parameter Note over Wrapper,Adapter: MindSpore optimizer state 通常在 optimizer 构造时已经存在 Train->>Wrapper: swap_optimizer(base_optimizer, config) Wrapper->>Adapter: 创建 MindFormersAdamWAdapter Adapter->>Adapter: 遍历优化器状态,注册SwapSlot Wrapper->>Runtime: offload_initial_slots(initial_slots) Runtime->>CPU: 为已有 device optimizer state 创建 pinned CPU mirror Runtime->>CPU: optimizer state D2H Runtime->>Runtime: 释放 device state storage Adapter->>Adapter: 获取 slots 创建 UpdateUnit Wrapper->>Runtime: partition(layout_units) Runtime->>Runtime: 按可 swap state 的 state_nbytes 均衡切分得到 layout_batches Wrapper->>Runtime: prepare_packed_host(layout_batches) loop 遍历每个 batch、dtype Runtime->>CPU: 创建该 batch、dtype 的连续 pinned CPU buffer Runtime->>CPU: 从 cpu_tensor 拷贝到 连续 pinned buffer 对应的 CPU view Runtime->>Runtime: 记录 slot.host_offset<br/>slot.cpu_tensor / slot.tensor 绑定 host view Runtime->>Runtime: 每个 batch 记录 PackedBatchPlan 信息 end Note over Train,Param: optimizer update Train->>Wrapper: construct(gradients) Wrapper->>Adapter: prepare_step(gradients) Adapter->>Adapter: 计算 lr / weight_decay,更新 global_step,预处理 gradients Adapter->>Adapter: 用最新 slots 构建 UpdateUnit[] Wrapper->>Runtime: partition(units) Runtime->>Runtime: 按可 swap state_nbytes 均衡切分 step batches Wrapper->>Runtime: prepare_packed_host(batches) alt packed layout signature 未变化 Runtime->>Runtime: 直接复用已有 batch/dtype pinned CPU buffers else packed slots 或 batches 发生变化 loop Runtime->>CPU: 按 batch、dtype 创建连续 pinned CPU buffer Runtime->>CPU: 将 CPU state 拷贝到对应 host view Runtime->>Runtime: 更新 host_offset 和 PackedBatchPlan end end Wrapper->>Runtime: run_pipeline -> _run_packed_pipeline par begin_packed_step(batches) Runtime->>Runtime: 读取 PackedBatchPlan Runtime->>Runtime: 统计在一 batch 中每种 dtype 所需最大的空间,<br>并用各 dtype 最大值计算一个布局,布局内各 dtype 区域按512对齐 Runtime->>A: 创建 staging arena A,按 dtype 创建 device views Runtime->>B: 创建 staging arena B,按 dtype 创建 device views end Note over Runtime,B: 以下以 batch i 使用 arena A 为例<br/>batch i+1 使用 B,batch i+2 再复用 A Runtime->>Copy: prefetch(batch 0, arena A) Copy->>A: 各 dtype packed region H2D Copy-->>Runtime: ready event Batch 0 Runtime->>Copy: prefetch(batch 1, arena B) Copy->>B: 各 dtype packed region H2D Copy-->>Runtime: ready event Batch 1 loop batch i,arena A Runtime->>Compute: wait batch i ready event alt i >= 2 Runtime->>Compute: wait batch i-2 offload event Runtime->>CPU: 将 batch i-2 slots 重新绑定CPU view end Runtime->>A: 将 batch i slots 绑定到 arena i%2 的 dtype view Runtime->>Adapter: step_batch(batch i) Adapter->>Compute: 调用 _run_fused_adamw_opt / _run_adamw_opt Adapter->>Compute: 通过 _slot_tensor() 传入 staging state view Compute->>Param: 原地更新 fp32/model parameter Compute->>A: 更新 staging 中的 exp_avg / exp_avg_sq<br/>以及可选 max_exp_avg_sq / master_param alt include_master_params=True Adapter->>Param: master parameter 转换并同步到 model parameter end Compute-->>Copy: 记录 update-complete event Copy->>A: 等待 update-complete event A->>CPU: Batch i packed D2H Copy->>A: 同一 copy stream 上执行 Batch i+2 packed H2D Copy-->>Runtime: offload&next-ready event end Note over Runtime,B: 最后两个 batch Runtime->>Compute: 等待最后两个 batch 的 offload event Runtime->>CPU: 将最后两个 batch 的 slots 重新绑定 CPU view Runtime->>A: 释放 arena A storage,storage.resize_(0) Runtime->>B: 释放 arena B storage,storage.resize_(0) Wrapper->>Adapter: finish_step() alt include_master_params=False Adapter->>Param: master parameter 同步到 model parameter end Wrapper-->>Train: construct() 返回
class SwapSlot: """One logical optimizer state tensor that may be swapped.""" name: str tensor: Any cpu_tensor: Optional[Any] = None storage_nbytes: int = 0 swappable: bool = True state: str = "device" event: Optional[Any] = None shape: tuple[int, ...] = () dtype: Optional[Any] = None device: Optional[Any] = None numel: int = 0 host_offset: int = 0 packed: bool = False logical_tensor: Optional[Any] = None
SwapSlot:一个 optimizer state 张量会被包装成一个 SwapSlot。
SwapSlot
name
tensor
cpu_tensor
storage_nbytes
swappable
state
pending
host
h2d
device
d2h
event
shape
dtype
numel
host_offset
packed
logical_tensor
class UpdateUnit: """Per-parameter optimizer update unit used by the pipeline runtime.""" adapter_index: int param: Any grad: Any slots: List[SwapSlot]
UpdateUnit:描述一次以“参数”为粒度的完整更新,把该参数的梯度和它依赖的全部 SwapSlot 组织在一起。运行时按 UpdateUnit 分批,保证更新一个参数所需要的所有状态同时在设备上。以 AdamW 的某个参数 P 为例:
UpdateUnit
P
UpdateUnit(P) ├── param: P ├── grad: P.grad └── slots ├── SwapSlot("exp_avg") ├── SwapSlot("exp_avg_sq") └── SwapSlot("max_exp_avg_sq") # AMSGrad 时
param
adapter_index
grad
slots
OptimizerSwapAdapter 定义了优化器更新的完整生命周期:
OptimizerSwapAdapter
matches()
validate()
prepare_step()
iter_update_units()
step_batch()
finish_step()
一次优化器更新的调用过程是:
prepare_step() | iter_update_units() | runtime.partition() | runtime.run_pipeline() runtime.prefetch() -> step_batch() -> runtime.offload() | finish_step()
继承自 class OptimizerSwapAdapter 定义torch优化器的 prepare_step(), step_batch() 等方法
class OptimizerSwapAdapter
TorchAdamBaseAdapter
TorchNativeAdamAdapter
TorchNativeAdamWAdapter
torch.optim.AdamWeightDecay
TorchHyperAdamWAdapter
继承自 class OptimizerSwapAdapter
MindSporeAdamBaseAdapter
MindSporeNativeAdamAdapter
mindspore.nn.Adam
MindSporeNativeAdamWAdapter
mindspore.nn.AdamWeightDecay
MindFormersAdamWAdapter
mindformers.pynative.optimizer.adamw.AdamW
主要方法:
# 普通逐 tensor 流水线: # 先 prefetch batch 0 # 对 batch n: # 等待 batch n 的 H2D # 等待 batch n-1 的 D2H,并释放其设备存储 # 提前发起 batch n+1 的 H2D # 计算 update batch n # 发起 batch n 的 D2H # 最后等待最后一批 D2H def run_pipeline( self, batches: Sequence[Sequence[UpdateUnit]], step_context: Any, step_batch: Callable[[List[UpdateUnit], Any], Any], ) -> List[Any]: """执行提前一批预取的普通流水线,并及时回收已完成的卸载任务。""" results = [] # 固化批次内容,确保后续异步操作始终引用同一组列表对象。 batch_lists = [list(batch) for batch in batches] if not batch_lists: return results # 所有批次满足 packed 条件时,改用双 staging buffer 流水线。 if self.supports_packed_pipeline(batch_lists): return self._run_packed_pipeline(batch_lists, step_context, step_batch) # 在进入循环前预取第 0 批,使首批计算可以尽快开始。 self.prefetch(batch_lists[0]) for index, batch_list in enumerate(batch_lists): # 当前批必须先完成 H2D 预取,优化器才能读取对应状态。 self.wait_prefetch(batch_list) previous_index = index - 1 if previous_index >= 0: # 扩大预取窗口前,先确认上一批 D2H 已完成并释放相关资源。 self.wait_offload(batch_lists[previous_index]) next_index = index + 1 if next_index < len(batch_lists): # 当前批计算期间,复制流可以并行预取下一批。 self.prefetch(batch_lists[next_index]) # 更新当前批,并刷新可能被适配器替换过的状态 tensor 引用。 results.append(step_batch(batch_list, step_context)) self.refresh_swappable_slots(batch_list) # 将更新后的状态异步卸载回主机,为后续批次腾出设备内存。 self.offload(batch_list) # 最后一批之后没有下一轮循环负责等待,因此在返回前显式收尾。 self.wait_offload(batch_lists[-1]) return results
# Packed 双缓冲流水线 # Copy Stream: # [H2D B0] # [H2D B1] # [D2H B0][H2D B2] # [D2H B1 # Compute Stream: # [Adam B0] # [Adam B1] # [Adam B2] def _run_packed_pipeline( self, batches: Sequence[List[UpdateUnit]], step_context: Any, step_batch: Callable[[List[UpdateUnit], Any], Any], ) -> List[Any]: """使用两个可复用的 staging buffer 执行 packed 状态更新流水线。 staging buffer 按批次下标奇偶固定复用。同一复制流链上,第 n 批执行 D2H 卸载后,第 n + 2 批才能复用相同 buffer 执行 H2D 预取;与此同时, 另一个 buffer 可供计算流更新相邻批次。 """ results = [] try: # 为本轮 optimizer step 创建 packed host 存储和两个 staging buffer。 self.begin_packed_step(batches) # 先填充最多两个 buffer,让计算阶段从第 0 批开始连续消费。 self.enqueue_packed_prefetch(0, 0) if len(batches) > 1: self.enqueue_packed_prefetch(1, 1) for batch_index, batch in enumerate(batches): staging_index = batch_index % 2 # 等待当前 buffer 的 H2D 完成,再允许计算流访问其中的数据。 self.wait_packed_prefetch(batch_index, staging_index) completed_index = batch_index - 2 if completed_index >= 0: # 当前批与前两批复用同一 buffer,复用前必须完成旧批次的 # D2H,并将卸载结果绑定回对应的逻辑状态。 self.wait_packed_offload(completed_index) self.finish_packed_offload(completed_index) # 把当前批的状态 slot 绑定到 staging buffer 中的对应视图。 self.activate_packed_batch(batch_index, staging_index) results.append(step_batch(batch, step_context)) self.refresh_swappable_slots(batch) # 在一条复制流链中先卸载当前批,再预取两批后的数据;二者 # 使用相同 buffer,按此顺序排队可避免状态被提前覆盖。 next_index = batch_index + 2 self.enqueue_packed_offload_prefetch( batch_index, next_index if next_index < len(batches) else None, staging_index, ) # 循环末尾最多还有两个已提交但未完成收尾的卸载任务。 drain_start = max(0, len(batches) - 2) for batch_index in range(drain_start, len(batches)): self.wait_packed_offload(batch_index) self.finish_packed_offload(batch_index) finally: # 即使更新或复制抛出异常,也要释放结果引用并销毁本轮临时存储。 self.release_packed_step_results(results) self.end_packed_step() return results
hyper_parallel/core/optimizer/swap_optimizer.py
swap_optimizer()
SwapOptimizerConfig
swap_times=16
min_numel=1024
include_master_params=False
swap_optimizer_base.py
PipelineSwapRuntime
TorchSwapOptimizer
torch.optim.Optimizer
__init__()
param_groups
zero_grad()
add_param_group()
hyper_parallel.core.optimizer.adamw.AdamW
numel >= min_numel
__call__()
construct()
include_master_params = True
use_lazy
use_offload
enable_cpu_offload
mindformers deepseekV3模型
Hidden size:1024 总参数量:25,239,552 FP32 参数内存:96.28 MiB 参数张量数量:816 AdamW moment 张量:1,632 AdamW state 内存:192.56 MiB
Hyper 提前预取了一个 batch 约 24 MiB 以拷贝与计算重叠。若要降低 Hyper 峰值,可以把 swap_times 从 8 增大到 16,预计额外分区内存降至约 12 MiB,但会增加拷贝事件。
tests/ut/core/optimizer/test_swap_optimizer.py
tests/ut/platform/mindspore/swap_optimizer/test_adapters.py
tests/ut/platform/torch/swap_optimizer/test_adapters.py
tests/mindspore/st/swap_optimizer/swap_optimizer.py
tests/mindspore/st/swap_optimizer/test_swap_optimizer.py
tests/mindspore/st/swap_optimizer/mf_adamw.py
tests/torch/swap_optimizer/swap_optimizer.py
test_native_adam_swap_optimizer_state_align:对比原生 Adam 与包装 swap optimizer 后训练多步的结果,验证参数、两个一阶/二阶矩状态及 beta 幂次一致,并确认 swap 可降低设备内存占用。
test_native_adam_nesterov_swap_optimizer_state_align:在 use_nesterov=True 的 Adam 场景下进行原生与 swap 训练对比,检查参数和优化器状态一致及内存下降。
test_native_adam_amsgrad_swap_optimizer_state_align:测试启用 AMSGrad 的 Adam swap,验证参数、moment、vhat 和 beta 状态对齐,同时确认 swap 节省设备内存。
test_native_adam_weight_decay_swap_optimizer_state_align:对比 AdamWeightDecay 原生版和 swap 版,检查参数及两个矩状态一致,并验证 swap 的内存优势。
test_mindformers_adamw_non_fused_swap_optimizer_state_align:测试 MindFormers 非 fused AdamW 的逐张量和 packed 两种 swap 模式,包含 fp32 master 参数交换,验证损失、参数、矩状态、全局步数和 master 参数一致。
test_mindformers_adamw_fused_swap_optimizer_state_align:测试 MindFormers fused AdamW 在逐张量和 packed swap 下的行为,验证训练结果、优化器状态及 fp32 master 参数与基线一致。
test_native_adam_fully_shard_swap_optimizer_state_align_worker:在 2×2 fully_shard 分布式模型上,对比原生 Adam 和 swap Adam,验证每步损失、最终本地参数分片及 Adam 状态一致。
test_mindformers_adamw_fully_shard_swap_optimizer_state_align_worker:在 fully_shard 场景测试 MindFormers AdamW 的非 packed 与 packed swap,验证损失、本地参数分片和优化器状态对齐。
test_native_adam_swap_optimizer_checkpoint_cpu_mirror_roundtrip:覆盖 Adam 和 AdamWeightDecay 的 checkpoint 保存/加载往返,验证可交换状态使用 CPU mirror、不可交换状态保留原值,并正确恢复到 swap optimizer。
test_native_adam_swap_optimizer_checkpoint_fresh_load_builds_slots:将 Adam checkpoint 加载到尚未训练的全新 swap optimizer,验证已有 slot 被复用、状态保持在 CPU,并以非严格模式完成加载。
test_mindformers_adamw_packed_swap_optimizer_checkpoint_roundtrip:保存并恢复 packed AdamW 的矩状态和 fp32 master 参数,验证 checkpoint 使用 CPU 副本恢复 packed slot,其他状态交由底层优化器加载。
test_torch_adam_swap_optimizer_parameter_align:对比 torch.optim.Adam 原生版与 swap 版多步训练,验证每步损失和最终参数一致,并确认 swap 降低峰值显存。
test_torch_adamw_swap_optimizer_parameter_align:测试原生 AdamW 与 swap AdamW 的训练一致性及显存降低效果。
test_torch_fused_adamw_swap_optimizer_parameter_align:针对 fused=True 的 AdamW,分别测试逐张量和 packed swap,验证结果对齐、状态卸载及显存下降。
test_torch_adamw_eager_state_swap_optimizer_parameter_align:先显式创建 AdamW 优化器状态再包装 swap,验证从第一步开始训练结果一致,并确认预先卸载状态可降低显存。
test_torch_adam_amsgrad_swap_optimizer_parameter_align:测试启用 AMSGrad 的 Adam swap,验证训练参数/损失一致及状态卸载和显存收益。
test_hyper_adamw_swap_optimizer_parameter_align:测试 HyperParallel AdamW 的 packed swap,验证参数和损失对齐、优化器 group step 正常递增、状态卸载且显存降低。
test_hyper_adamw_amsgrad_swap_optimizer_parameter_align:测试 HyperParallel AdamW 开启 AMSGrad 并使用 packed swap,验证训练一致性、packed 存储和显存收益。
test_torch_adam_swap_optimizer_multi_param_group_align:使用不同学习率、权重衰减和 betas 的多个参数组测试 Adam swap,验证参数/损失一致、pipeline 按参数组正确拆批及显存降低。
test_fully_shard_adamw_mixed_precision_swap_optimizer_parameter_align:在 fully_shard 混合精度策略下,分别测试 PyTorch AdamW 和 HyperParallel AdamW 的 swap,验证损失、本地参数分片和优化器状态一致。
test_fully_shard_optimizer_swap_adamw_4card_parameter_align:在四卡 fully_shard 上分别验证非 packed 与 packed AdamW swap,检查结果和 checkpoint 恢复一致、存储模式正确及显存低于原生 AdamW。
test_torch_adam_swap_optimizer_checkpoint_host_state:覆盖 PyTorch Adam、AdamW 和 HyperParallel AdamW 的 host-resident checkpoint,验证保存无需完整回迁到设备,加载后状态继续驻留主机,并在下一步训练时按需预取。
数值对齐 同一模型、同一随机种子下,开启/关闭 swap_optimizer 后,loss、参数更新结果、optimizer state(exp_avg/exp_avg_sq/max_exp_avg_sq)应与基线一致。
分布式/FSDP 场景 重点验FSDP场景和TP/EP/CP/PP场景下 loss 和 optimizer state 能与基线对齐。
packed_swap 与普通 swap 分别验证 packed_swap=True/False 两条路径,关注默认值是否符合设计:Torch 默认开启 packed,MindSpore 默认关闭,但 MindFormers AdamW 默认应开启 packed。
master params / state_keys / min_numel 配置项验证: include_master_params 是否只在该支持的优化器上生效; state_keys 只 swap 指定 state 时是否正确; min_numel 调大后小 tensor 是否不再被 swap。
异常/不支持场景拦截 不支持场景报错是否符合预期,例如: Torch 的 closure、foreach/capturable/differentiable; MindSpore 的 use_lazy/use_offload/use_parallel; 原生 MindSpore Adam/AdamWeightDecay 配 packed_swap=True 应被明确拒绝。
内存与状态迁移正确性 验证 step 前后 optimizer state 是否真的发生 device/host 迁移,保存 checkpoint 前 CPU mirror 是否已同步,避免“功能能跑但实际没释放显存/NPU 内存”。
1. 基本信息
swap/checkpoint/optimizer2. 背景
目前的训练中Adam/AdamW优化器状态会长期驻留在设备显存中,对大模型而言,显存容量成为更早、更明显的瓶颈。
Swap optimizer 的核心思路是:在两次 optimizer update 之间,将 Adam/AdamW 的大状态张量保存在 pinned CPU 内存中;只在执行 update 时把当前批次状态搬到设备,借助独立拷贝流尽量把下一批 H2D、当前批计算和上一批 D2H 重叠,从而降低 optimizer step 期间的设备峰值内存。
3. 目标和非目标
3.1 目标
3.2 非目标
4. 对外接口
4.1 接口定义
4.1.1 swap_optimizer(optimizer, config=None)
包装已有 Adam/AdamW optimizer,返回当前框架对应的 swap optimizer。之后继续按原框架方式使用:
optimizer.step()、optimizer.zero_grad()optimizer(gradients)约束:
torch.optim.Adam、torch.optim.AdamW、HyperParallelAdamW。nn.Adam、nn.AdamWeightDecay、MindFormers PyNativeAdamW。config=None使用默认配置。Torch 侧和 HyperParallelAdamW优化器默认packed_swap=True,MindSpore 侧nn.Adam、nn.AdamWeightDecay优化器默认packed_swap=False。4.1.2 SwapOptimizerConfig(...)
swap_times16> 0state_keysNone其他还支持
exp_avg、exp_avg_sq、max_exp_avg_sq、master_parammin_numel1024>= 0include_master_paramsFalsepacked_swapTrue/FalseTrue使用 A/B packed buffer,False使用逐 tensor swap其他说明:
state_keys=None表示自动选择已有优化器状态。max_exp_avg_sq。min_numel的状态才会真正 swap;其他状态继续常驻设备。packed_swap=True/False。AdamW支持packed_swap=True,其他只支持False。4.2 使用示例
from hyper_parallel.core.optimizer import SwapOptimizerConfig, swap_optimizer optimizer = swap_optimizer( base_optimizer, SwapOptimizerConfig( swap_times=16, min_numel=1024, packed_swap=True, ), )4.3 对 base_optimizer 的限制
1. torch.optim.Adam/AdamW
class torch.optim.Adam(params, lr=0.001, betas=(0.9, 0.999), eps=1e-08, weight_decay=0, amsgrad=False, *, foreach=None, maximize=False, capturable=False, differentiable=False, fused=None, decoupled_weight_decay=False) class torch.optim.AdamW(params, lr=0.001, betas=(0.9, 0.999), eps=1e-08, weight_decay=0.01, amsgrad=False, *, maximize=False, foreach=None, capturable=False, differentiable=False, fused=None) Optimizer.step(closure: None = None) → None[source]foreach=False、fused=False、capturable=False、differentiable=False,否则swap_optimizer包装时直接 ValueError。(torch-npu 2.10: aten::_fused_adam_is not currently supported on the NPU backend)foreach=False、capturable=False、differentiable=False,否则swap_optimizer包装时直接 ValueError。step()不支持closure,配置直接 ValueError。2. HyperParallel AdamW
hyper_parallel.core.optimizer.adamw( params: List[torch.Tensor], grads: List[torch.Tensor], exp_avgs: List[torch.Tensor], exp_avg_sqs: List[torch.Tensor], max_exp_avg_sqs: List[torch.Tensor], step: int, *, amsgrad: bool, beta1: float, beta2: float, lr: float, weight_decay: float, eps: float, maximize: bool ) Optimizer.step(closure: None = None)step()不支持closure,配置直接 ValueError。3. mindspore.nn.Adam/AdamWeightDecay
class mindspore.nn.Adam(params, learning_rate=1e-3, beta1=0.9, beta2=0.999, eps=1e-8, use_locking=False, use_nesterov=False, weight_decay=0.0, loss_scale=1.0, use_amsgrad=False, **kwargs) class mindspore.nn.AdamWeightDecay(params, learning_rate=1e-3, beta1=0.9, beta2=0.999, eps=1e-6, weight_decay=0.0)use_lazy=False、use_offload=False,否则swap_optimizer包装时直接 ValueError。4. Mindformers AdamW
class AdamW( params, learning_rate=1e-3, betas=(0.9, 0.999), eps=1e-8, weight_decay=0.0, enable_cpu_offload=False, enable_fused_opt=False, use_fused=False, **kwargs )enable_cpu_offload=False,否则swap_optimizer包装时直接 ValueError。include_master_params配置为 True 时,swap 低精度参数对应的优化器自有的 FP32 master parameter。5. 方案设计
5.1 总体流程
per tensor 模式
以每个 optimizer state 为独立单位执行搬运。

优点:适配范围广、实现灵活。
缺点:小块拷贝和显存分配次数较多,调度开销较大。
per tensor swap 掩盖关系:
packed 模式
将相同 dtype 的多个状态 tensor 打包进连续的 pinned CPU buffer,每step使用两个可复用的 device staging buffer。

优点:把大量小拷贝合并成连续拷贝,减少 allocation 和调度开销,并通过 A/B 双缓冲重叠状态传输与参数更新。
缺点:使用有限制
packed swap 掩盖关系:
torch 侧 packed swap 内存结构:

mindspore 侧 packed swap 内存结构:

5.2 时序参考
5.2.1 Torch
5.2.1.1 per tensor swap
sequenceDiagram autonumber participant Train as Training Loop participant Wrapper as TorchSwapOptimizer participant Adapter as TorchAdamBaseAdapter participant Runtime as PipelineSwapRuntime participant Copy as Copy Stream participant CPU as CPU Pinned Memory participant Compute as Compute Stream participant Param as Model Parameter Note over Wrapper,CPU: 初始化阶段,config.packed_swap=False Train->>Wrapper: swap_optimizer(base_optimizer, config) Wrapper->>Adapter: 创建 optimizer adapter Adapter->>Runtime: 遍历 param_groups(处理目前已存在state,如checkpoint加载),<br>注册 exp_avg / exp_avg_sq 等 SwapSlot Wrapper->>Runtime: offload_initial_slots(initial_slots) Runtime->>CPU: 为所有已存在的 optimizer state 创建/获取 pinned CPU mirror Runtime->>CPU: D2H copy optimizer state Runtime->>Runtime: 释放 device state storage Note over Train,Param: 一次 optimizer.step() Train->>Wrapper: optimizer.step() Wrapper->>Adapter: prepare_step() loop 遍历有梯度的 parameter Adapter->>Adapter: 懒初始化 optimizer state: <br>创建 pinned CPU mirror -> 创建 device tensor 以构建 SwapSlot 后随即 resize(0) Adapter->>Adapter: 为当前参数的每个 state 找到/创建 SwapSlot Adapter->>Adapter: 获取 slots 创建 UpdateUnit(param, grad, slots) <br>(UpdateUnit 包含一个 param 的所有要做swap 的 optimizer states) end Adapter-->>Wrapper: UpdateUnit 列表 Wrapper->>Runtime: partition(units) Runtime->>Runtime: 按可 swap 的优化器状态的 state_nbytes 近似均衡切分 batch Wrapper->>Runtime: run_pipeline Note over Runtime,CPU: 首批预取 Runtime->>Copy: prefetch(batch 0) Runtime->>Runtime: 恢复 batch 0 每个 state tensor 的 device storage CPU->>Copy: H2D param 优化器状态(exp_avg、exp_avg_sq...) Copy->>Copy: 记录 batch 0 H2D ready event loop batch i Runtime->>Compute: wait batch i H2D complete event Runtime->>Runtime: wait (batch i-1) D2H complete event,释放 device state storage Runtime->>Copy: prefetch(batch i+1) <br> 恢复 device storage,做 H2D Runtime->>Adapter: step_batch(batch i) Adapter->>Compute: functional Adam/AdamW Compute->>Param: 原地更新 model parameter Compute->>Compute: 原地更新 exp_avg / exp_avg_sq Compute->>Copy: 记录 update-complete event Runtime->>Copy: offload(batch i):wait update-complete event, 做 D2H Copy->>Copy: 记录 D2H complete event end Runtime->>Runtime: wait last D2H event,释放 device state storage Runtime-->>Wrapper: pipeline 完成 Wrapper-->>Train: step() 返回 Note over Param,CPU: step 结束状态 Note over Param: Model parameter 已在设备上原地更新 Note over CPU: swappable moments 保存在 CPU pinned memory5.2.1.2 packed swap
sequenceDiagram autonumber participant Train as Training Loop participant Wrapper as TorchSwapOptimizer participant Adapter as Adam Adapter participant Runtime as Packed Runtime participant CPU as CPU Pinned Buffer participant Copy as Copy Stream participant A as Device Arena A participant B as Device Arena B participant Compute as Compute Stream participant Param as Model Parameter Note over Wrapper,CPU: 初始化阶段,config.packed_swap=True Train->>Wrapper: swap_optimizer(base_optimizer, config) Wrapper->>Adapter: 创建 optimizer adapter Adapter->>Runtime: 遍历 param_groups(处理目前已存在state,如checkpoint加载),<br>注册 exp_avg / exp_avg_sq 等 SwapSlot Wrapper->>Runtime: ofload_initial_slots(initial_slots) Runtime->>CPU: 为所有已存在的 optimizer state 创建/获取 pinned CPU mirror Runtime->>CPU: D2H copy optimizer state Runtime->>Runtime: 释放 device state storage Wrapper->>Runtime: prepare_packed_host(initial_slots) loop 遍历 dtype Runtime->>CPU: 按 dtype 创建连续 pinned CPU buffer Runtime->>CPU: 将 cpu_tensor 复制到连续 pinned CPU buffer 中对应的一段view Runtime->>Runtime: 注册每个 slot 的 host_offset,<br>设置 slot 的 cpu_tensor 为对应 host_view end Note over Train,Param: optimizer.step() Train->>Wrapper: step() Wrapper->>Adapter: prepare_step() loop 遍历有梯度的 parameter Adapter->>Adapter: 懒初始化 optimizer state: 创建 SwapSlot Adapter->>Adapter: 登记现有的 SwapSlot end Adapter->>Runtime: prepare_packed_host() alt packed_slots 完全没变 Runtime->>Runtime: 直接复用旧的连续 CPU buffer else packed_slots 有变化 loop Runtime->>CPU: 按 dtype 创建连续 pinned CPU buffer Runtime->>CPU: 将 cpu_tensor 或 tensor 复制到连续 pinned CPU buffer 中对应的一段view Runtime->>Runtime: 注册每个 slot 的 host_offset,<br>设置 slot 的 cpu_tensor 为对应 host_view end end Adapter->>Adapter: 获取 slots 创建 UpdateUnit(...) (UpdateUnit 包含一个 param 的所有要做swap 的 optimizer states) Adapter-->>Wrapper: 返回 UpdateUnit[] Wrapper->>Runtime: partition(units) Runtime->>Runtime: 按可 swap 的优化器状态的 state_nbytes 均衡切分 batch Wrapper->>Runtime: run_pipeline -> _run_packed_pipeline par begin_packed_step(batches) Runtime->>Runtime: 遍历 batches ,收集每 batch 需要 swap 的 packed slots,每 batch slots 按 dtype 分组 Runtime->>Runtime: 然后 batch 中 slots 根据 dtype 集中排列,组合成连续 region -> 为每个 batch 创建 PackedBatchPlan 记录 region 信息 Runtime->>Runtime: 统计在一 batch 中每种 dtype 所需最大的空间,<br>并用各 dtype 最大值计算一个布局,布局内各 dtype 区域按512对齐 Runtime->>A: 创建 staging arena A,按 dtype 创建 device views Runtime->>B: 创建 staging arena B,按 dtype 创建 device views end Note over Runtime,A: 准备 batch 0 Runtime->>Copy: prefetch(batch 0, arena A) Copy->>A: Batch 0 连续 region H2D Copy-->>Runtime: ready event Batch 0 Note over Runtime,B: 预取 batch 1 Runtime->>Copy: prefetch(batch 1, arena B) Copy->>B: Batch 1 连续 region H2D Copy-->>Runtime: ready event Batch1 loop 例如 batch i 使用 arena A Runtime->>Compute: wait batch i alt i >= 2 Runtime->>Compute: wait batch i-2 offload event Runtime->>CPU: 将 batch i-2 slots 重新绑定CPU view end Runtime->>A: 将 Batch i slots 绑定到 arena i%2 Runtime->>Adapter: step_batch(batch i) Adapter->>Compute: functional Adam/AdamW Compute->>Param: 原地更新模型参数 Compute->>Compute: 更新 staging 中的 exp_avg/exp_avg_sq... Compute-->>Copy: 记录 update-complete event Copy->>A: 等待 update-complete A->>CPU: batch i D2H CPU->>A: 同一 copy stream 上执行 batch i+2 H2D Copy-->>Runtime: offload&next-ready event end Note over Runtime,B: 最后2个 batch Copy->>Runtime: 等待最后的 offload event Runtime->>CPU: 将 slots 重新绑定 CPU view Note over Runtime,B: step mo末尾 Runtime->>A: 释放 arena A storage Runtime->>B: 释放 arena B storage Wrapper-->>Train: step() 返回5.2.2 Mindspore
5.2.2.1 per tensor swap
sequenceDiagram autonumber participant Train as Training Loop participant Wrapper as MindSporeSwapOptimizer participant Adapter as MindSpore Adam Adapter participant BaseOpt as Base Optimizer participant Runtime as MindSporeSwapRuntime participant Copy as Copy Stream participant CPU as CPU Pinned Memory participant Compute as Compute Stream participant Param as Model / Master Parameter Note over Wrapper,Adapter: MindSpore optimizer state 通常在 optimizer 构造时已经存在 Train->>Wrapper: swap_optimizer(base_optimizer, config) Wrapper->>Adapter: 创建 adapter Adapter->>BaseOpt: 遍历已创建的 moment1/moment2/vhat<br>或 exp_avg/exp_avg_sq/fp32_params 优化器状态,注册SwapSlot <br>(ms 的 optimizer 在构造时就把 optimizer state 创建好了) Wrapper->>Runtime: offload_initial_slots(initial_slots) Runtime->>CPU: 为所有已存在的 optimizer state 创建/获取 pinned CPU mirror Runtime->>CPU: D2H copy optimizer state Runtime->>Runtime: 释放 device state storage Note over Train,Param: 一次 optimizer update Train->>Wrapper: construct(gradients) Wrapper->>Adapter: prepare_step(gradients) Adapter->>BaseOpt: 校验/预处理梯度,获取 lr、weight decay<br>更新 global_step / beta powers loop 遍历有梯度的 parameter Adapter->>Adapter: 获取 state slots(初始化时已创建) Adapter->>Adapter: 创建 UpdateUnit(param, grad, selected slots) end Adapter-->>Wrapper: UpdateUnit 列表 Wrapper->>Runtime: partition(units) Runtime->>Runtime: 按 swappable slot.storage_nbytes<br>保持原顺序近似均衡切分 Wrapper->>Runtime: run_pipeline() Note over Runtime,CPU: 首批预取 Runtime->>Copy: prefetch(batch 0) Runtime->>Runtime: 恢复 batch 0 每个 state tensor 的 device storage CPU->>Copy: H2D param 优化器状态 Copy->>Copy: 记录 batch 0 H2D ready event loop batch i Runtime->>Compute: wait batch i H2D complete event Runtime->>Runtime: wait (batch i-1) D2H complete event,释放 device state storage Runtime->>Copy: prefetch(batch i+1) <br> 恢复 device storage,做 H2D Runtime->>Adapter: step_batch(batch i) alt mindspore.nn.Adam Adapter->>Compute: 每个 unit 调用 _apply_adam / AMSGrad opt Compute->>Param: 原地更新 model param 和 moments else mindspore.nn.AdamWeightDecay Adapter->>Compute: 每个 unit 调用 fused_opt Compute->>Param: 原地更新 model param 和 moments else MindFormers AdamW Adapter->>Compute: _run_adamw_opt / _run_fused_adamw_opt Compute->>Param: 原地更新 model param 和 moments opt include_master_params=True Adapter->>Param: 当前 batch fp32 master param 同步到 model param <br>(因为马上就会offload fp32 master param) end end Runtime->>Compute: record update-complete event Runtime->>Runtime: refresh_swappable_slots<br>(ms 有些优化器状态在第一次 step 后才会建立有效的 device storage) Runtime->>Copy: offload(batch i):wait update-complete event,做D2H Copy->>Copy: record D2H-complete event end Runtime->>Runtime: wait last D2H event,释放 device state storage Wrapper->>Adapter: finish_step() opt MindFormers 且 include_master_params=False Adapter->>Param: 将全部 fp32 main params 同步到 model params end Wrapper-->>Train: tuple(batch results) Note over Param,CPU: 更新结束 Note over CPU: swappable moments 由 pinned CPU mirror 持有 Note over Param: model param 已更新5.2.2.2 packed swap
sequenceDiagram autonumber participant Train as Training Loop participant Wrapper as MindSporeSwapOptimizer participant Adapter as MindFormers AdamW Adapter participant Runtime as Packed Runtime participant CPU as CPU Pinned Buffer participant Copy as Copy Stream participant A as Device Staging Arena A participant B as Device Staging Arena B participant Compute as Compute Stream participant Param as Model Parameter Note over Wrapper,Adapter: MindSpore optimizer state 通常在 optimizer 构造时已经存在 Train->>Wrapper: swap_optimizer(base_optimizer, config) Wrapper->>Adapter: 创建 MindFormersAdamWAdapter Adapter->>Adapter: 遍历优化器状态,注册SwapSlot Wrapper->>Runtime: offload_initial_slots(initial_slots) Runtime->>CPU: 为已有 device optimizer state 创建 pinned CPU mirror Runtime->>CPU: optimizer state D2H Runtime->>Runtime: 释放 device state storage Adapter->>Adapter: 获取 slots 创建 UpdateUnit Wrapper->>Runtime: partition(layout_units) Runtime->>Runtime: 按可 swap state 的 state_nbytes 均衡切分得到 layout_batches Wrapper->>Runtime: prepare_packed_host(layout_batches) loop 遍历每个 batch、dtype Runtime->>CPU: 创建该 batch、dtype 的连续 pinned CPU buffer Runtime->>CPU: 从 cpu_tensor 拷贝到 连续 pinned buffer 对应的 CPU view Runtime->>Runtime: 记录 slot.host_offset<br/>slot.cpu_tensor / slot.tensor 绑定 host view Runtime->>Runtime: 每个 batch 记录 PackedBatchPlan 信息 end Note over Train,Param: optimizer update Train->>Wrapper: construct(gradients) Wrapper->>Adapter: prepare_step(gradients) Adapter->>Adapter: 计算 lr / weight_decay,更新 global_step,预处理 gradients Adapter->>Adapter: 用最新 slots 构建 UpdateUnit[] Wrapper->>Runtime: partition(units) Runtime->>Runtime: 按可 swap state_nbytes 均衡切分 step batches Wrapper->>Runtime: prepare_packed_host(batches) alt packed layout signature 未变化 Runtime->>Runtime: 直接复用已有 batch/dtype pinned CPU buffers else packed slots 或 batches 发生变化 loop Runtime->>CPU: 按 batch、dtype 创建连续 pinned CPU buffer Runtime->>CPU: 将 CPU state 拷贝到对应 host view Runtime->>Runtime: 更新 host_offset 和 PackedBatchPlan end end Wrapper->>Runtime: run_pipeline -> _run_packed_pipeline par begin_packed_step(batches) Runtime->>Runtime: 读取 PackedBatchPlan Runtime->>Runtime: 统计在一 batch 中每种 dtype 所需最大的空间,<br>并用各 dtype 最大值计算一个布局,布局内各 dtype 区域按512对齐 Runtime->>A: 创建 staging arena A,按 dtype 创建 device views Runtime->>B: 创建 staging arena B,按 dtype 创建 device views end Note over Runtime,B: 以下以 batch i 使用 arena A 为例<br/>batch i+1 使用 B,batch i+2 再复用 A Runtime->>Copy: prefetch(batch 0, arena A) Copy->>A: 各 dtype packed region H2D Copy-->>Runtime: ready event Batch 0 Runtime->>Copy: prefetch(batch 1, arena B) Copy->>B: 各 dtype packed region H2D Copy-->>Runtime: ready event Batch 1 loop batch i,arena A Runtime->>Compute: wait batch i ready event alt i >= 2 Runtime->>Compute: wait batch i-2 offload event Runtime->>CPU: 将 batch i-2 slots 重新绑定CPU view end Runtime->>A: 将 batch i slots 绑定到 arena i%2 的 dtype view Runtime->>Adapter: step_batch(batch i) Adapter->>Compute: 调用 _run_fused_adamw_opt / _run_adamw_opt Adapter->>Compute: 通过 _slot_tensor() 传入 staging state view Compute->>Param: 原地更新 fp32/model parameter Compute->>A: 更新 staging 中的 exp_avg / exp_avg_sq<br/>以及可选 max_exp_avg_sq / master_param alt include_master_params=True Adapter->>Param: master parameter 转换并同步到 model parameter end Compute-->>Copy: 记录 update-complete event Copy->>A: 等待 update-complete event A->>CPU: Batch i packed D2H Copy->>A: 同一 copy stream 上执行 Batch i+2 packed H2D Copy-->>Runtime: offload&next-ready event end Note over Runtime,B: 最后两个 batch Runtime->>Compute: 等待最后两个 batch 的 offload event Runtime->>CPU: 将最后两个 batch 的 slots 重新绑定 CPU view Runtime->>A: 释放 arena A storage,storage.resize_(0) Runtime->>B: 释放 arena B storage,storage.resize_(0) Wrapper->>Adapter: finish_step() alt include_master_params=False Adapter->>Param: master parameter 同步到 model parameter end Wrapper-->>Train: construct() 返回5.4 关键逻辑
5.4.1 optimizer states 数据管理
5.4.1.1 class SwapSlot
class SwapSlot: """One logical optimizer state tensor that may be swapped.""" name: str tensor: Any cpu_tensor: Optional[Any] = None storage_nbytes: int = 0 swappable: bool = True state: str = "device" event: Optional[Any] = None shape: tuple[int, ...] = () dtype: Optional[Any] = None device: Optional[Any] = None numel: int = 0 host_offset: int = 0 packed: bool = False logical_tensor: Optional[Any] = NoneSwapSlot:一个 optimizer state 张量会被包装成一个SwapSlot。nameexp_avg、exp_avg_sq、max_exp_avg_sq、master_param。tensor逐 tensor 模式下非优化器更新时通常是设备张量空壳,优化器更新时是拥有有效设备 storage 的原始 Tensor/DTensor;
packed 模式下会在 CPU view 和device staging view 之间重新绑定。
cpu_tensorstorage_nbytesswappableFalse。statepending、host、h2d、device和d2h。pending是 lazy init 时 optimizer state 尚未物化的状态eventshapedtypedevicenumelhost_offsetpackedlogical_tensorfully_shard 场景下实际搬运的是 local tensor,但 optimizer 仍需要看到原来的 DTensor 语义。
由于 packed swap 模式会不断替换 slot.tensor 的指向对象, 所以另需一个 logical_tensor 保留原 DTensor 包装器,
否则优化器更新时 slot.tensor 会丢失 mesh,placements 等信息。
5.4.1.2 class UpdateUnit
class UpdateUnit: """Per-parameter optimizer update unit used by the pipeline runtime.""" adapter_index: int param: Any grad: Any slots: List[SwapSlot]UpdateUnit:描述一次以“参数”为粒度的完整更新,把该参数的梯度和它依赖的全部SwapSlot组织在一起。运行时按UpdateUnit分批,保证更新一个参数所需要的所有状态同时在设备上。以 AdamW 的某个参数P为例:paramadapter_index在 MindSpore 侧是该参数在优化器参数列表的下标,用于索引该参数的 moment1,moment2,lr 等。
在 Torch 侧是指该参数属于的 optimizer parameter group 下标,用于索引对应的lr、eps、weight_decay 等配置。
gradslotsSwapSlot。5.4.2 优化器语义适配
5.4.2.1 class OptimizerSwapAdapter
OptimizerSwapAdapter定义了优化器更新的完整生命周期:matches():用于匹配支持本 optimizer 的 optimizer Adapter 是否支持这个 optimizer。validate():拒绝无法支持的配置。prepare_step():执行一次性的 step 准备,收集梯度、学习率、global step 等。iter_update_units():返回本轮所有参数更新单元。step_batch():执行一批参数的 Adam/AdamW 更新。finish_step():完成优化器更新后处理,例如 master parameter 的同步。一次优化器更新的调用过程是:
5.4.2.2 class TorchAdamBaseAdapter
继承自
class OptimizerSwapAdapter定义torch优化器的 prepare_step(), step_batch() 等方法
继承自
TorchAdamBaseAdapterTorchNativeAdamAdaptertorch.optim.AdamTorchNativeAdamWAdaptertorch.optim.AdamWeightDecayTorchHyperAdamWAdapterAdamW5.4.2.3 class MindSporeAdamBaseAdapter
继承自
class OptimizerSwapAdapter继承自
MindSporeAdamBaseAdapterMindSporeNativeAdamAdaptermindspore.nn.Adam并根据本优化器特点定义 prepare_step(), step_batch() 等方法
MindSporeNativeAdamWAdaptermindspore.nn.AdamWeightDecayMindFormersAdamWAdaptermindformers.pynative.optimizer.adamw.AdamW5.4.3 优化器状态搬运流水线
主要方法:
# 普通逐 tensor 流水线: # 先 prefetch batch 0 # 对 batch n: # 等待 batch n 的 H2D # 等待 batch n-1 的 D2H,并释放其设备存储 # 提前发起 batch n+1 的 H2D # 计算 update batch n # 发起 batch n 的 D2H # 最后等待最后一批 D2H def run_pipeline( self, batches: Sequence[Sequence[UpdateUnit]], step_context: Any, step_batch: Callable[[List[UpdateUnit], Any], Any], ) -> List[Any]: """执行提前一批预取的普通流水线,并及时回收已完成的卸载任务。""" results = [] # 固化批次内容,确保后续异步操作始终引用同一组列表对象。 batch_lists = [list(batch) for batch in batches] if not batch_lists: return results # 所有批次满足 packed 条件时,改用双 staging buffer 流水线。 if self.supports_packed_pipeline(batch_lists): return self._run_packed_pipeline(batch_lists, step_context, step_batch) # 在进入循环前预取第 0 批,使首批计算可以尽快开始。 self.prefetch(batch_lists[0]) for index, batch_list in enumerate(batch_lists): # 当前批必须先完成 H2D 预取,优化器才能读取对应状态。 self.wait_prefetch(batch_list) previous_index = index - 1 if previous_index >= 0: # 扩大预取窗口前,先确认上一批 D2H 已完成并释放相关资源。 self.wait_offload(batch_lists[previous_index]) next_index = index + 1 if next_index < len(batch_lists): # 当前批计算期间,复制流可以并行预取下一批。 self.prefetch(batch_lists[next_index]) # 更新当前批,并刷新可能被适配器替换过的状态 tensor 引用。 results.append(step_batch(batch_list, step_context)) self.refresh_swappable_slots(batch_list) # 将更新后的状态异步卸载回主机,为后续批次腾出设备内存。 self.offload(batch_list) # 最后一批之后没有下一轮循环负责等待,因此在返回前显式收尾。 self.wait_offload(batch_lists[-1]) return results# Packed 双缓冲流水线 # Copy Stream: # [H2D B0] # [H2D B1] # [D2H B0][H2D B2] # [D2H B1 # Compute Stream: # [Adam B0] # [Adam B1] # [Adam B2] def _run_packed_pipeline( self, batches: Sequence[List[UpdateUnit]], step_context: Any, step_batch: Callable[[List[UpdateUnit], Any], Any], ) -> List[Any]: """使用两个可复用的 staging buffer 执行 packed 状态更新流水线。 staging buffer 按批次下标奇偶固定复用。同一复制流链上,第 n 批执行 D2H 卸载后,第 n + 2 批才能复用相同 buffer 执行 H2D 预取;与此同时, 另一个 buffer 可供计算流更新相邻批次。 """ results = [] try: # 为本轮 optimizer step 创建 packed host 存储和两个 staging buffer。 self.begin_packed_step(batches) # 先填充最多两个 buffer,让计算阶段从第 0 批开始连续消费。 self.enqueue_packed_prefetch(0, 0) if len(batches) > 1: self.enqueue_packed_prefetch(1, 1) for batch_index, batch in enumerate(batches): staging_index = batch_index % 2 # 等待当前 buffer 的 H2D 完成,再允许计算流访问其中的数据。 self.wait_packed_prefetch(batch_index, staging_index) completed_index = batch_index - 2 if completed_index >= 0: # 当前批与前两批复用同一 buffer,复用前必须完成旧批次的 # D2H,并将卸载结果绑定回对应的逻辑状态。 self.wait_packed_offload(completed_index) self.finish_packed_offload(completed_index) # 把当前批的状态 slot 绑定到 staging buffer 中的对应视图。 self.activate_packed_batch(batch_index, staging_index) results.append(step_batch(batch, step_context)) self.refresh_swappable_slots(batch) # 在一条复制流链中先卸载当前批,再预取两批后的数据;二者 # 使用相同 buffer,按此顺序排队可避免状态被提前覆盖。 next_index = batch_index + 2 self.enqueue_packed_offload_prefetch( batch_index, next_index if next_index < len(batches) else None, staging_index, ) # 循环末尾最多还有两个已提交但未完成收尾的卸载任务。 drain_start = max(0, len(batches) - 2) for batch_index in range(drain_start, len(batches)): self.wait_packed_offload(batch_index) self.finish_packed_offload(batch_index) finally: # 即使更新或复制抛出异常,也要释放结果引用并销毁本轮临时存储。 self.release_packed_step_results(results) self.end_packed_step() return results5.5 代码改动点
公共 API
hyper_parallel/core/optimizer/swap_optimizer.py增加统一入口swap_optimizer()。hyper_parallel/core/optimizer/swap_optimizer.py增加swap optimizer配置SwapOptimizerConfig: 可配参数:swap_times=16,min_numel=1024,state_keys=None,include_master_params=False(仅在使用MF的adamw优化器时生效),packed_swap=True/False。swap_optimizer_base.py抽象SwapSlot、UpdateUnit、OptimizerSwapAdapter和PipelineSwapRuntime,把状态识别、计算逻辑与搬运调度解耦。Torch 侧关键改动
TorchSwapOptimizer继承torch.optim.Optimizer,重写__init__();param_groups、state、zero_grad()、add_param_group()等仍托付给被包装优化器。torch.optim.Adam、torch.optim.AdamW和hyper_parallel.core.optimizer.adamw.AdamW优化器。foreach=False、capturable=False、differentiable=False的 functional 路径,以便只更新当前 swap batch。numel >= min_numel,且 tensor 独占完整 storage;其他状态继续常驻设备。MindSpore 侧关键改动
__call__(),construct()包装;其他属性继续委托给原 optimizer。mindspore.nn.Adam、mindspore.nn.AdamWeightDecay,以及mindformers.pynative.optimizer.adamw.AdamW优化器。include_master_params = True。use_lazy/use_offload,以及 MindFormers 自带的enable_cpu_offload,防止两套 offload 机制冲突。offload。
numel >= min_numel,且 tensor 独占完整 storage;其他状态继续常驻设备。
开始。
packed_swap=True。MindSpore 原生 Adam 和AdamWeightDecay 默认
packed_swap=False。6. 收益与劣化
mindspore
mindformers deepseekV3模型
torch
Hidden size:1024
总参数量:25,239,552
FP32 参数内存:96.28 MiB
参数张量数量:816
AdamW moment 张量:1,632
AdamW state 内存:192.56 MiB
Hyper 提前预取了一个 batch 约 24 MiB 以拷贝与计算重叠。若要降低 Hyper 峰值,可以把 swap_times 从 8 增大到 16,预计额外分区内存降至约 12 MiB,但会增加拷贝事件。
7. 验证设计
7.1 组件交互验证
7.2 上库用例设计
7.2.1 测试文件列表
tests/ut/core/optimizer/test_swap_optimizer.pytests/ut/platform/mindspore/swap_optimizer/test_adapters.pytests/ut/platform/torch/swap_optimizer/test_adapters.pytests/mindspore/st/swap_optimizer/swap_optimizer.pytests/mindspore/st/swap_optimizer/test_swap_optimizer.pytests/mindspore/st/swap_optimizer/mf_adamw.pytests/torch/swap_optimizer/swap_optimizer.pytests/torch/swap_optimizer/swap_optimizer.py7.2.2 测试场景
MindSpore 场景
test_native_adam_swap_optimizer_state_align:对比原生 Adam 与包装 swap optimizer 后训练多步的结果,验证参数、两个一阶/二阶矩状态及 beta 幂次一致,并确认 swap 可降低设备内存占用。
test_native_adam_nesterov_swap_optimizer_state_align:在 use_nesterov=True 的 Adam 场景下进行原生与 swap 训练对比,检查参数和优化器状态一致及内存下降。
test_native_adam_amsgrad_swap_optimizer_state_align:测试启用 AMSGrad 的 Adam swap,验证参数、moment、vhat 和 beta 状态对齐,同时确认 swap 节省设备内存。
test_native_adam_weight_decay_swap_optimizer_state_align:对比 AdamWeightDecay 原生版和 swap 版,检查参数及两个矩状态一致,并验证 swap 的内存优势。
test_mindformers_adamw_non_fused_swap_optimizer_state_align:测试 MindFormers 非 fused AdamW 的逐张量和 packed 两种 swap 模式,包含 fp32 master 参数交换,验证损失、参数、矩状态、全局步数和 master 参数一致。
test_mindformers_adamw_fused_swap_optimizer_state_align:测试 MindFormers fused AdamW 在逐张量和 packed swap 下的行为,验证训练结果、优化器状态及 fp32 master 参数与基线一致。
test_native_adam_fully_shard_swap_optimizer_state_align_worker:在 2×2 fully_shard 分布式模型上,对比原生 Adam 和 swap Adam,验证每步损失、最终本地参数分片及 Adam 状态一致。
test_mindformers_adamw_fully_shard_swap_optimizer_state_align_worker:在 fully_shard 场景测试 MindFormers AdamW 的非 packed 与 packed swap,验证损失、本地参数分片和优化器状态对齐。
test_native_adam_swap_optimizer_checkpoint_cpu_mirror_roundtrip:覆盖 Adam 和 AdamWeightDecay 的 checkpoint 保存/加载往返,验证可交换状态使用 CPU mirror、不可交换状态保留原值,并正确恢复到 swap optimizer。
test_native_adam_swap_optimizer_checkpoint_fresh_load_builds_slots:将 Adam checkpoint 加载到尚未训练的全新 swap optimizer,验证已有 slot 被复用、状态保持在 CPU,并以非严格模式完成加载。
test_mindformers_adamw_packed_swap_optimizer_checkpoint_roundtrip:保存并恢复 packed AdamW 的矩状态和 fp32 master 参数,验证 checkpoint 使用 CPU 副本恢复 packed slot,其他状态交由底层优化器加载。
Torch 场景
test_torch_adam_swap_optimizer_parameter_align:对比 torch.optim.Adam 原生版与 swap 版多步训练,验证每步损失和最终参数一致,并确认 swap 降低峰值显存。
test_torch_adamw_swap_optimizer_parameter_align:测试原生 AdamW 与 swap AdamW 的训练一致性及显存降低效果。
test_torch_fused_adamw_swap_optimizer_parameter_align:针对 fused=True 的 AdamW,分别测试逐张量和 packed swap,验证结果对齐、状态卸载及显存下降。
test_torch_adamw_eager_state_swap_optimizer_parameter_align:先显式创建 AdamW 优化器状态再包装 swap,验证从第一步开始训练结果一致,并确认预先卸载状态可降低显存。
test_torch_adam_amsgrad_swap_optimizer_parameter_align:测试启用 AMSGrad 的 Adam swap,验证训练参数/损失一致及状态卸载和显存收益。
test_hyper_adamw_swap_optimizer_parameter_align:测试 HyperParallel AdamW 的 packed swap,验证参数和损失对齐、优化器 group step 正常递增、状态卸载且显存降低。
test_hyper_adamw_amsgrad_swap_optimizer_parameter_align:测试 HyperParallel AdamW 开启 AMSGrad 并使用 packed swap,验证训练一致性、packed 存储和显存收益。
test_torch_adam_swap_optimizer_multi_param_group_align:使用不同学习率、权重衰减和 betas 的多个参数组测试 Adam swap,验证参数/损失一致、pipeline 按参数组正确拆批及显存降低。
test_fully_shard_adamw_mixed_precision_swap_optimizer_parameter_align:在 fully_shard 混合精度策略下,分别测试 PyTorch AdamW 和 HyperParallel AdamW 的 swap,验证损失、本地参数分片和优化器状态一致。
test_fully_shard_optimizer_swap_adamw_4card_parameter_align:在四卡 fully_shard 上分别验证非 packed 与 packed AdamW swap,检查结果和 checkpoint 恢复一致、存储模式正确及显存低于原生 AdamW。
test_torch_adam_swap_optimizer_checkpoint_host_state:覆盖 PyTorch Adam、AdamW 和 HyperParallel AdamW 的 host-resident checkpoint,验证保存无需完整回迁到设备,加载后状态继续驻留主机,并在下一步训练时按需预取。
8. 验收checklist
数值对齐
同一模型、同一随机种子下,开启/关闭 swap_optimizer 后,loss、参数更新结果、optimizer state(exp_avg/exp_avg_sq/max_exp_avg_sq)应与基线一致。
分布式/FSDP 场景
重点验FSDP场景和TP/EP/CP/PP场景下 loss 和 optimizer state 能与基线对齐。
packed_swap 与普通 swap
分别验证 packed_swap=True/False 两条路径,关注默认值是否符合设计:Torch 默认开启 packed,MindSpore 默认关闭,但 MindFormers AdamW 默认应开启 packed。
master params / state_keys / min_numel 配置项验证:
include_master_params 是否只在该支持的优化器上生效;
state_keys 只 swap 指定 state 时是否正确;
min_numel 调大后小 tensor 是否不再被 swap。
异常/不支持场景拦截
不支持场景报错是否符合预期,例如:
Torch 的 closure、foreach/capturable/differentiable;
MindSpore 的 use_lazy/use_offload/use_parallel;
原生 MindSpore Adam/AdamWeightDecay 配 packed_swap=True 应被明确拒绝。
内存与状态迁移正确性
验证 step 前后 optimizer state 是否真的发生 device/host 迁移,保存 checkpoint 前 CPU mirror 是否已同步,避免“功能能跑但实际没释放显存/NPU 内存”。