mindspore-ai / mindspore-ai/hyper-parallel

[RFC]: swap optimizer

Open
#653 0 comments 0 reactions 0 assignees View on GitHub

Nobody has claimed this yet.

Dominant language
Python
Stars
53
Forks
63
Avg merge
23h 45m
Merged PRs (30d)
63

Description

1. 基本信息

项目 内容
作者 宋佳琪
相关模块 swap / checkpoint / optimizer
相关 issue / PR https://gitcode.com/mindspore/hyper-parallel/pull/982
适用后端 PT + MS

2. 背景

目前的训练中Adam/AdamW优化器状态会长期驻留在设备显存中,对大模型而言,显存容量成为更早、更明显的瓶颈。
Swap optimizer 的核心思路是:在两次 optimizer update 之间,将 Adam/AdamW 的大状态张量保存在 pinned CPU 内存中;只在执行 update 时把当前批次状态搬到设备,借助独立拷贝流尽量把下一批 H2D、当前批计算和上一批 D2H 重叠,从而降低 optimizer step 期间的设备峰值内存。

3. 目标和非目标

3.1 目标
1. 支持优化器范围:torch: 原生 Adam、AdamW、HP自研 AdamW;mindspore: 原生Adam、AdamWeightDecay、MF自研 AdamW;
2. 支持swap optimizer特性与FSDP,TP/PP/EP/CP,重计算,swap activation等特性叠加使用。
3.2 非目标
1. 对已支持优化器的使用有限制,以便以batch为粒度更新;
2. 本期不支持 muon 优化器,待后期补全。

4. 对外接口

4.1 接口定义
4.1.1 swap_optimizer(optimizer, config=None)

包装已有 Adam/AdamW optimizer,返回当前框架对应的 swap optimizer。之后继续按原框架方式使用:

  • Torch:optimizer.step()optimizer.zero_grad()
  • MindSpore:optimizer(gradients)

约束:

  • Torch 支持 torch.optim.Adamtorch.optim.AdamW、HyperParallel AdamW
  • MindSpore 支持 nn.Adamnn.AdamWeightDecay、MindFormers PyNative AdamW
  • 目前不支持其他 optimizer。
  • config=None 使用默认配置。Torch 侧和 HyperParallel AdamW 优化器默认 packed_swap=True,MindSpore 侧 nn.Adamnn.AdamWeightDecay 优化器默认 packed_swap=False
4.1.2 SwapOptimizerConfig(...)
参数 默认值 含义与约束
swap_times 16 optimizer state 的目标分批数,必须 > 0
state_keys None 指定要 swap 的 optimizer state, 配为 None 则 swap 所有符合要求的 states;
其他还支持 exp_avgexp_avg_sqmax_exp_avg_sqmaster_param
min_numel 1024 状态 tensor 元素数至少达到该值才参与 swap,必须 >= 0
include_master_params False 是否 swap 优化器自有的 FP32 master parameter,仅对 MindFormers AdamW 有效
packed_swap True/False True 使用 A/B packed buffer,False 使用逐 tensor swap

其他说明:

  • state_keys=None 表示自动选择已有优化器状态。
  • 显式指定 optimizer 不具备的状态会报错,例如未启用 AMSGrad 时指定 max_exp_avg_sq
  • 只有浮点、连续、非 CPU、独占完整 storage 且达到 min_numel 的状态才会真正 swap;其他状态继续常驻设备。
  • Torch 支持 packed_swap=True/False
  • MindSpore 只有 HyperParallel 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]
  • Adam 创建时必须设置 foreach=Falsefused=Falsecapturable=Falsedifferentiable=False ,否则swap_optimizer 包装时直接 ValueError。(torch-npu 2.10: aten::_fused_adam_ is not currently supported on the NPU backend)
  • AdamW 创建时必须设置 foreach=Falsecapturable=Falsedifferentiable=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)
  • nn.Adam 优化器创建时必须设置 use_lazy=Falseuse_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。
  • SwapOptimizerConfig 中 include_master_params 配置为 True 时,swap 低精度参数对应的优化器自有的 FP32 master parameter。

5. 方案设计

5.1 总体流程
per tensor 模式

以每个 optimizer state 为独立单位执行搬运。
优点:适配范围广、实现灵活。
缺点:小块拷贝和显存分配次数较多,调度开销较大。
per tensor swap 掩盖关系:
per_tensor_pipeline_timeline.png

packed 模式

将相同 dtype 的多个状态 tensor 打包进连续的 pinned CPU buffer,每step使用两个可复用的 device staging buffer。
优点:把大量小拷贝合并成连续拷贝,减少 allocation 和调度开销,并通过 A/B 双缓冲重叠状态传输与参数更新。
缺点:使用有限制
packed swap 掩盖关系:
image.png

torch 侧 packed swap 内存结构:
pt_packed_内存结构.png

mindspore 侧 packed swap 内存结构:
ms_packed_内存结构.png

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 memory
5.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] = None

SwapSlot:一个 optimizer state 张量会被包装成一个 SwapSlot

属性 用途
name optimizer state 的逻辑名称,例如 exp_avgexp_avg_sqmax_exp_avg_sqmaster_param
tensor 当前暴露给 optimizer 的活动张量。
逐 tensor 模式下非优化器更新时通常是设备张量空壳,优化器更新时是拥有有效设备 storage 的原始 Tensor/DTensor;
packed 模式下会在 CPU view 和device staging view 之间重新绑定。
cpu_tensor CPU 上的副本,是 pinned memory,以支持异步 H2D/D2H;packed 模式下是大块 host buffer 的一个 view。
storage_nbytes 恢复该张量设备 storage 所需要的字节数,也是 pipeline 分批时的依据。
swappable 标记该 slot 是否参与搬运。小张量、CPU 张量、非连续张量或共享 storage 等为 False
state 当前生命周期状态,包括 pendinghosth2ddeviced2hpending 是 lazy init 时 optimizer state 尚未物化的状态
event 最近一次异步 H2D/D2H 对应的 stream event。计算流据此等待拷贝完成。
shape 逻辑状态或 DTensor local tensor 的形状。packed staging 中从一维 buffer 重建 view 时使用。
dtype 数据类型。用于按 dtype 构建 packed host buffer、device staging 区域和检查 packed pipeline 是否可用。
device optimizer update 应发生的目标设备。用于 packed pipeline 的同设备校验。
numel 这个 slot 的元素数。用于计算 offset、packed buffer 切片和 staging view 构造。
host_offset 此 slot 在对应 dtype 的 packed host buffer 中的元素偏移,不是字节偏移。用于定位 H2D/D2H 区间。
packed 是否是 packed swap 候选,但不代表最终一定走 packed;运行时还会检查。
logical_tensor 服务于 fully_shard + packed swap 场景:
fully_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 为例:

UpdateUnit(P)
├── param: P
├── grad: P.grad
└── slots
    ├── SwapSlot("exp_avg")
    ├── SwapSlot("exp_avg_sq")
    └── SwapSlot("max_exp_avg_sq")  # AMSGrad 时
属性 功能说明
param 被 optimizer 更新的参数。
adapter_index 适配不同框架下优化器状态的更新方式:
在 MindSpore 侧是该参数在优化器参数列表的下标,用于索引该参数的 moment1,moment2,lr 等。
在 Torch 侧是指该参数属于的 optimizer parameter group 下标,用于索引对应的lr、eps、weight_decay 等配置。
grad 当前参数的梯度。
slots 该参数更新所依赖的全部 SwapSlot
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 的同步。

一次优化器更新的调用过程是:

prepare_step()
        |
iter_update_units()
        |
runtime.partition()
        |
runtime.run_pipeline()
  	runtime.prefetch() -> step_batch() -> runtime.offload()
        |
finish_step()
5.4.2.2 class TorchAdamBaseAdapter

继承自 class OptimizerSwapAdapter
定义torch优化器的 prepare_step(), step_batch() 等方法

Adapter
继承自TorchAdamBaseAdapter
对应优化器 用途
TorchNativeAdamAdapter torch.optim.Adam 只用于识别是否本 optimizer
TorchNativeAdamWAdapter torch.optim.AdamWeightDecay 同上
TorchHyperAdamWAdapter Hyper Parallel 自己的 AdamW 同上
5.4.2.3 class MindSporeAdamBaseAdapter

继承自 class OptimizerSwapAdapter

Adapter
继承自MindSporeAdamBaseAdapter
对应优化器 用途
MindSporeNativeAdamAdapter mindspore.nn.Adam 用于识别是否本 optimizer,
并根据本优化器特点定义 prepare_step(), step_batch() 等方法
MindSporeNativeAdamWAdapter mindspore.nn.AdamWeightDecay 同上
MindFormersAdamWAdapter mindformers.pynative.optimizer.adamw.AdamW 同上
5.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 results
5.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 抽象 SwapSlotUpdateUnitOptimizerSwapAdapterPipelineSwapRuntime,把状态识别、计算逻辑与搬运调度解耦。
  • 实现 packed/per-tensor 两种流水的统一调度接口。
Torch 侧关键改动
  1. 优化器适配
  • TorchSwapOptimizer 继承 torch.optim.Optimizer,重写__init__()param_groupsstatezero_grad()add_param_group() 等仍托付给被包装优化器。
  • 支持 torch.optim.Adamtorch.optim.AdamWhyper_parallel.core.optimizer.adamw.AdamW 优化器。
  • 将原 optimizer 的整步更新拆为按 batch 调用 Torch functional Adam/AdamW 更新。
  • Torch 原生 Adam/AdamW 强制走 foreach=Falsecapturable=Falsedifferentiable=False 的 functional 路径,以便只更新当前 swap batch。
  1. 状态生命周期与搬运
  • 兼容 Torch 优化器状态的 lazy init。
  • swap optimizer 创建前已经物化 的 optimizer state 会被发现并立即 offload。
  • 仅 swap 满足以下条件的状态:float、contiguous、非 CPU、numel >= min_numel,且 tensor 独占完整 storage;其他状态继续常驻设备。
  • Per tensor swap:使用 pinned CPU mirror、独立 copy stream 上进行 H2D/D2H;H2D 前设备 storage resize 按字节数恢复,D2H event 完成后将设备 storage resize 为 0。
  1. Packed swap
  • host 侧,按 dtype 把所有 swappable 状态放入持久 pinned host buffer,并记录每个 SwapSlot 的 host offset。
  • 每 step 只创建/恢复两块设备 buffer;不同 dtype 区域按 512 字节对齐,并从 buffer 构造每个 optimizer state tensor 的 view。
  • 支持 DTensor:实际 storage 管理作用于 local tensor,更新期间把 logical wrapper 绑定到 staging view,结束后重新绑定 CPU mirror。
  • FSDP 场景在 packed step 开始前会做一次设备同步,防止 FSDP 仍有 stream 有任务未完成。
MindSpore 侧关键改动
  1. 优化器适配
  • MindSporeSwapOptimizer提供__call__(), construct() 包装;其他属性继续委托给原 optimizer。
  • 支持 mindspore.nn.Adammindspore.nn.AdamWeightDecay,以及mindformers.pynative.optimizer.adamw.AdamW 优化器。
  • 按 batch 调用原优化器的优化器更新方法(_run_fused_adamw_opt、fused_op...)。
  • MindFormers AdamW 可选择将 optimizer-owned fp32 master parameter 一并 swap,随后同步回低精度模型参数,需要配置 include_master_params = True
  • 拒绝 nn.Adam/AdamWeightDecay 的 use_lazy/use_offload,以及 MindFormers 自带的 enable_cpu_offload,防止两套 offload 机制冲突。
  1. 状态生命周期与搬运
  • MindSpore optimizer 状态通常在构造时已存在,因此 swap optimizer 初始化时直接构建 slot 并
    offload。
  • 仅 swap 满足以下条件的状态:float、contiguous、非 CPU、numel >= min_numel
    且 tensor 独占完整 storage;其他状态继续常驻设备。
  • Per tensor swap:使用 pinned CPU mirror、独立 copy stream 上进行 H2D/D2H;H2D 前设备 storage resize 按字节数恢复,D2H event 完成后将设备 storage resize 为 0。
  1. Packed Swap
  • MindSpore packed host buffer 按 batch 和 dtype 持久化;每个传输 buffer 都从 storage offset 0
    开始。
  • 两块 NPU staging buffer 按 512 字节对齐,并从 buffer 构造每个 optimizer state tensor 的 view。每个 batch/dtype 直接在pinned host buffer 与 staging view 之间执行一次 H2D/D2H。
  • 当前只有 MindFormers AdamW 支持 packed_swap=True。MindSpore 原生 Adam 和
    AdamWeightDecay 默认 packed_swap=False

6. 收益与劣化

mindspore

mindformers deepseekV3模型

指标 基线 per tensor packed 测试结果
step time 3528ms 3584ms 3534ms 劣化 1.59% \ 0.17%
optimizer time 366ms 441ms 392ms 劣化 20.49% \ 7.10%
peak memory 20.26 G 19.40 G 19.40 G 降低 4.24%
loss 对齐 可对齐
torch

Hidden size:1024
总参数量:25,239,552
FP32 参数内存:96.28 MiB
参数张量数量:816
AdamW moment 张量:1,632
AdamW state 内存:192.56 MiB

Mode Step mean Optimizer mean Peak NPU Idle NPU Host state Mem Earn
Native fused 123.18 ms 42.75 ms 402.93 MiB 386.93 MiB 0 /
竞品 664.46 ms 580.09 ms 233.51 MiB 193.36 MiB 192.56 MiB 42%
Hyper per-tensor 195.92 ms 115.49 ms 258.05 MiB 193.76 MiB 192.56 MiB 36%
Hyper packed 186.69 ms 105.69 ms 257.90 MiB 193.76 MiB 192.56 MiB 36%

Hyper 提前预取了一个 batch 约 24 MiB 以拷贝与计算重叠。若要降低 Hyper 峰值,可以把 swap_times 从 8 增大到 16,预计额外分区内存降至约 12 MiB,但会增加拷贝事件。

7. 验证设计

7.1 组件交互验证
组合 是否验证 验证方式 通过标准
本特性 + FSDP 和单卡/旧实现比 loss;检查 state_dict 功能正确,无精度问题。
本特性 + TP / PP / EP / CP 训练若干 step;检查 shape、通信、调度 功能正确,无精度问题。
本特性 + 重计算 固定 seed 比 loss 和显存 功能正确,无精度问题。
本特性 + swap/offload 检查搬运、峰值显存、性能劣化 功能正确,无精度问题。
本特性 + checkpoint save/load 后继续训练 功能正确,无精度问题。
本特性 + optimizer 检查 optimizer state state 无遗漏。
本特性 + PT/MS 双后端同配置运行或验证降级 差异符合文档。
7.2 上库用例设计
7.2.1 测试文件列表
文件 说明
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 mindspore 集成测试 Worker
tests/mindspore/st/swap_optimizer/test_swap_optimizer.py mindspore 集成测试 入口
tests/mindspore/st/swap_optimizer/mf_adamw.py mindformers 实现的 adamw,用于测试
tests/torch/swap_optimizer/swap_optimizer.py torch 集成测试 Worker
tests/torch/swap_optimizer/swap_optimizer.py torch 集成测试 入口
7.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

  1. 数值对齐
    同一模型、同一随机种子下,开启/关闭 swap_optimizer 后,loss、参数更新结果、optimizer state(exp_avg/exp_avg_sq/max_exp_avg_sq)应与基线一致。

  2. 分布式/FSDP 场景
    重点验FSDP场景和TP/EP/CP/PP场景下 loss 和 optimizer state 能与基线对齐。

  3. packed_swap 与普通 swap
    分别验证 packed_swap=True/False 两条路径,关注默认值是否符合设计:Torch 默认开启 packed,MindSpore 默认关闭,但 MindFormers AdamW 默认应开启 packed。

  4. master params / state_keys / min_numel 配置项验证:
    include_master_params 是否只在该支持的优化器上生效;
    state_keys 只 swap 指定 state 时是否正确;
    min_numel 调大后小 tensor 是否不再被 swap。

  5. 异常/不支持场景拦截
    不支持场景报错是否符合预期,例如:
    Torch 的 closure、foreach/capturable/differentiable;
    MindSpore 的 use_lazy/use_offload/use_parallel;
    原生 MindSpore Adam/AdamWeightDecay 配 packed_swap=True 应被明确拒绝。

  6. 内存与状态迁移正确性
    验证 step 前后 optimizer state 是否真的发生 device/host 迁移,保存 checkpoint 前 CPU mirror 是否已同步,避免“功能能跑但实际没释放显存/NPU 内存”。

schema_version: 1
source: gitcode
gitcode_repo: mindspore/hyper-parallel
gitcode_issue: 313
source_url: https://gitcode.com/mindspore/hyper-parallel/issues/313

Contributor guide

No contributing guide indexed for this repository

First steps

  1. Read the whole issue, then the project's contributing guide.
  2. Comment on the issue to say you are picking it up — it saves two people doing the same work.
  3. Fork the repository and make your change on a branch.
  4. Open a pull request that references the issue number.

Research direction

Start by reading the swap, checkpoint, and optimizer modules and the related PR linked in the RFC. Trace the proposed swap_optimizer and SwapOptimizerConfig entry points across Torch and MindSpore support. Done means the listed Adam/AdamW combinations work with the documented parallelism and recomputation features, while unsupported configurations fail as specified.

Written by the indexing model from the issue text.

Assessment

Tech stack
python, pytorch
Domain
backend, machine-learning, performance
Issue type
Feature
Difficulty
5/5
Estimated time
Over a week
Activity status
Stale
Clarity
Needs clarification
Newbie friendliness
25/100

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.