mindspore-ai / mindspore-ai/hyper-parallel

[RFC] HyperOffload:自动异步激活值卸载

Open
#693 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. 背景与动机 (Background & Motivation)

HyperParallel 框架中训练大规模 Transformer / LLM 时,前向传播产生的中间激活值 (activations) 是长序列 (long-context) 场景下设备内存 (Device Memory) 不足并触发 OOM 的核心原因之一。随着序列长度、批量大小和模型层数的同步增长,激活值占用的显存往往呈线性甚至超线性膨胀,严重制约了可训练模型的规模上限。

项目中虽已存在基于 activation_checkpointswap 机制(通过 CheckpointPolicy.MUST_SWAPSwapManager 手动包裹特定算子/模块),但该方案在工程实践中暴露出以下痛点:

  1. 侵入性高:用户需要深入理解网络结构,手动挑选并包裹待卸载的算子或子模块。模型迭代时,这些包裹点需要重新调优,维护成本随模型复杂度急剧上升。
  2. 粒度较粗:以模块(Module)为单位进行 Swap,无法细粒度控制单个算子的激活,难以将传输与计算充分重叠。

因此,我们需要一种零侵入、细粒度、自动异步卸载的新机制:HyperOffload


2. 目标与非目标 (Goals & Non-Goals)

2.1 目标 (Goals)
  • 零侵入接入:用户无需修改模型代码,仅通过 with OffloadSession(config): 上下文即可开启激活值自动卸载。
  • 自动预算感知调度:根据峰值显存预算自动决定哪些激活何时卸载、何时预取回来。
  • 细粒度驻留控制:字节级存储追踪,独立的 D2H 异步拷贝、设备释放、H2D 预取、主机释放,在专用拷贝流上执行。
  • 数学等价性:卸载与预取过程不改变前向 loss 与反向梯度,保证与基线完全一致的数值结果。
  • 可扩展架构:分层设计(API / IR / Execution / Planning / Runtime),便于后续接入更多 Planner、更多后端(MindSpore / 其他加速器)以及参数/优化器卸载。
2.2 非目标 (Non-Goals)
  • 本次不涉及参数卸载或优化器状态卸载HyperOffload 聚焦中间激活值卸载,模型参数与优化器状态仍由 FSDP / HSDP / ZeRO 等既有模块管理。
  • 本次首版仅支持 PyTorch 后端:MindSpore 等后端因 dispatch 机制差异,将在后续 RFC 中单独设计适配层。
  • 不替代 activation checkpoint 本身HyperOffload 可与 checkpoint / recompute 叠加使用;它不是重计算,而是显存换带宽的互补技术。
  • 不保证所有动态控制流自动处理:对无法被 TorchDispatchMode 精确追踪的分支/循环控制流,提供 @skip_offload 逃生舱,由用户显式标注。

3. 设计概览 (Design Overview)

HyperOffload 采用 "先追踪、后规划、再重放" 的两阶段执行范式:

┌──────────────────────────────────────────────────────────────────────────────┐
│                              用户训练脚本                                      │
│  config = OffloadConfig(max_resident_activation_mb=512)                      │
│  session = OffloadSession(config)                                            │
│  with session:                                                               │
│      loss = model(x)           # Warmup Step: 记录 Trace + 在线驱逐            │
│      loss.backward()                                                         │
│                                                                              │
│  # __exit__ 时自动完成:                                                       │
│  #   1. GreedyResidencyPlanner 生成 ResidencySchedule                         │
│  #   2. WarmupExecutor -> ReplayExecutor                                     │
│                                                                              │
│  with session:                                                               │
│      loss = model(x)           # Replay Step: 按 Schedule 异步卸载/预取        │
│      loss.backward()                                                         │
└──────────────────────────────────────────────────────────────────────────────┘

                      ┌───────────────────────┐
                      │    OffloadSession     │
                      │  (context manager)    │
                      └───────────┬───────────┘
                                  │
              ┌───────────────────┼───────────────────┐
              │                   │                   │
              ▼                   ▼                   ▼
    ┌─────────────────┐ ┌─────────────────┐ ┌─────────────────┐
    │   API Layer     │ │  Execution Layer│ │  Planning Layer │
    │  OffloadConfig  │ │ WarmupExecutor  │ │GreedyResidency  │
    │  skip_offload   │ │ ReplayExecutor  │ │    Planner      │
    └─────────────────┘ └─────────────────┘ └─────────────────┘
              │                   │                   │
              │                   ▼                   │
              │         ┌───────────────────┐         │
              │         │   IR Layer        │         │
              │         │ ActivationTrace   │◄────────┘
              │         │ ResidencySchedule │
              │         │    OpGuide        │
              │         └───────────────────┘
              │                   │
              ▼                   ▼
    ┌─────────────────────────────────────────────────────┐
    │                  Runtime Layer                      │
    │  ResidencyManager + PinnedMemoryPool + BandwidthEst │
    └─────────────────────────────────────────────────────┘

核心设计哲学:

  • Trace:首个 step 作为 Warmup,通过 TorchDispatchMode 拦截所有 Tensor 操作,记录每个激活值 storage 的创建、读写、销毁时序。
  • Plan:Warmup 退出时,GreedyResidencyPlanner 根据全局显存预算和 access interval,生成一个离线驻留调度表 (ResidencySchedule)。
  • Replay:后续 step 切换为 ReplayExecutor,严格按调度表执行 COPY_D2HRELEASE_DEVICECOPY_H2DRELEASE_HOST,并通过独立拷贝流与计算流同步实现重叠。
  • ResidencyResidencyManagerPinnedMemoryPool 负责物理 buffer 的分配、异步拷贝、事件同步与回收。

4. 详细设计 (Detailed Design)

4.1 API 层 (hyper_parallel/auto_parallel/hyper_offload/api/)
OffloadConfig

配置入口,当前暴露三个关键字段:

@dataclass
class OffloadConfig:
    max_resident_activation_mb: int = 1024      # 设备端驻留激活值上限 (MiB)
    max_offload_activation_mb: int = 65536      # 固定主机内存池上限 (MiB),默认 64 GiB
    planner: ResidencyPlanner | None = None     # 可插拔规划器,默认 GreedyResidencyPlanner
OffloadSession

继承上下文管理器语义,内部维护两阶段 Executor:

  • 第一次进入 __enter__ 时处于 warmup 模式,使用 WarmupExecutor
  • __exit__ 时等待所有异步传输完成,调用 WarmupExecutor.finish() 得到 ActivationTraceOpGuide,再由 planner.build(trace) 生成 ResidencySchedule,随后切换为 ReplayExecutor
  • 之后再次进入 session 即进入 replay 模式,按调度表执行。

OffloadSession 通过 contextvars.ContextVar 维护当前活跃 session,供 @skip_offload 查询。

skip_offload

装饰器/透明 API,用于标记一段代码为 Opaque Region(虚拟 op)。被装饰函数内部的算子不再被逐个追踪,而是作为单个虚拟 op 记录。适用于动态控制流、第三方库函数或用户自定义算子。

@skip_offload
def my_custom_block(x):
    # 内部大量细粒度算子不会被单独追踪
    return some_dynamic_logic(x)
4.2 IR 层 (hyper_parallel/auto_parallel/hyper_offload/ir/)
ActivationTrace

Warmup 阶段的完整记录,包含:

  • ops: list[TraceOp]:按执行顺序排列的 op,每个 op 记录算子名、耗时、输入/输出 storage 访问(READ / WRITE)。
  • storage_sizes: dict[sid, bytes]:每个 storage ID 的字节大小。
  • retained_sids: set[int]:Warmup 结束时仍有 ShadowTensor 存活的 storage(通常是反向仍需要的激活值)。
  • memory_limit_bytes:设备端显存预算。
  • d2h_bandwidth_gbps / h2d_bandwidth_gbps:实测或默认的传输带宽,供后续 planner 参考。
ResidencySchedule

Planner 输出,按 op 索引维护:

  • pre_actions[op_id]:op 执行前需要完成的动作(目前主要是 COPY_H2D 预取)。
  • post_actions[op_id]:op 执行后需要完成的动作(COPY_D2HRELEASE_DEVICERELEASE_HOST)。
OpGuide

ReplayExecutor 使用的预消化结构,避免在热路径上重复解析 ActivationTrace。每个 op 保存输出叶子数量与 leaf_index -> storage_id 的绑定关系,用于快速 ShadowTensor 包裹与结构校验。

4.3 执行层 (hyper_parallel/auto_parallel/hyper_offload/execution/)
BaseExecutor

定义统一的生命周期钩子:

  • on_op_begin(func, args, kwargs):op 开始前,缓存 func/args/kwargs,递增 op 索引。
  • on_op_end(result):op 结束后,记录 trace/执行 post-actions,并将输出 Tensor 替换为 ShadowTensor
  • dispatch(func, args, kwargs):标准分发模板。若处于 Opaque Region 则直接调用 func;否则走 on_op_begin -> func -> on_op_end
  • execute_opaque_op(...):将一段函数包装为虚拟 op,并通过 OpaqueRegionStart / OpaqueRegionEnd 两个 autograd.Function 维持反向图连续性。
WarmupExecutor
  • 在线记录 ActivationTraceOpGuide
  • on_op_begin 时执行 在线贪心驱逐:若当前驻留字节数超过 memory_limit_bytes,则按 "最早产生 op 优先、同 op 大小大者优先" 的策略选择 victim,调用 copy_d2h + release_device
  • on_op_end 时:
    • 通过 ActivationTracker 识别本 op 新产生的 activation storage;
    • 检测 mutable alias(func._schema.is_mutable)以标记 write access;
    • 生成 TraceOpOpGuide
  • finish() 结束时进行带宽 profile,并返回 trace + guide。
ReplayExecutor
  • 进入 replay 后,严格按 ResidencySchedule 执行 pre/post actions。
  • on_op_begin 执行 COPY_H2D 预取。
  • on_op_end 执行 COPY_D2HRELEASE_DEVICERELEASE_HOST
  • 校验输出叶子数量与 OpGuide.output_leaf_count 一致,确保模型结构与 Warmup 一致。
ShadowTensor

一个 torch.Tensor 子类(通过 _make_wrapper_subclass 构造),本身不持有 device 数据,而是持有对 PhysicalBuffer 的引用。每次被调度时通过 resolve()PhysicalBuffer.device_storage() 重新构造视图:

  • 若数据已在 device,直接返回视图;
  • 若数据仅在 host,则同步 demand-page 回 device( Warmup / 异常回退路径)。

ShadowTensor 不缓存长期 device 引用,因此不会阻止底层 storage 被释放。

4.4 规划层 (hyper_parallel/auto_parallel/hyper_offload/planning/)
GreedyResidencyPlanner

核心离线算法,目标是在满足 memory_limit_bytes 的前提下,选择一组 "access gap" 进行卸载,使得传输次数与暴露的传输延迟最小化。

算法步骤

  1. ActivationTrace 按 storage ID 分组,得到每个 storage 的访问序列(按 op_id 排序)。
  2. 计算每个 storage 在每个 op 上的驻留字节贡献,得到 resident_bytes[op_id]
  3. 对每对相邻访问 (release_start, end) 构造候选 _EvictionCandidate
    • distance = end.op_id - release_start.op_id(gap 长度)
    • copy_startrelease_start 之前最近的一次 WRITE 访问(确保 host 副本不会过期)
    • copy_start.op_id < end.op_id - 1 则视为有效候选。
  4. 将候选按 (distance, size, -release_start.op_id) 降序排序:优先卸载 "距离下次使用最远、尺寸最大" 的 activation。
  5. 依次选择候选,只要该候选覆盖的任一 op 当前 resident_bytes[op_id] 仍超过预算,就将其加入调度表,并扣减相应 op 的驻留字节。
  6. 为每个被选中的候选生成:
    • COPY_D2H @ copy_start.op_id
    • RELEASE_DEVICE @ release_start.op_id
    • COPY_H2D @ end.op_id
  7. 对非 retained storage,在最后一次访问后追加 RELEASE_DEVICE;若该 storage 曾被卸载,再追加 RELEASE_HOST 以归还 host 内存。

复杂度:设 op 数为 S,storage 数为 N,候选数为 M。排序 O(M log M),模拟 O(M * S),对典型训练图足够高效;若后续需要,可引入线段树 / 差分数组优化到 O(M log S)

4.5 运行时层 (hyper_parallel/auto_parallel/hyper_offload/runtime/)
ResidencyManager

物理驻留控制器,维护 storage_id -> PhysicalBuffer 映射:

  • bind(sid, tensor):将新 tensor 的 UntypedStorage 注册到 PhysicalBuffer,返回 buffer。
  • copy_d2h(sid):在独立 copy_stream 上发起异步 D2H 拷贝;对跨设备场景(copy stream 与 tensor 不同设备)回退同步拷贝。
  • copy_h2d(sid):在独立 copy_stream 上发起异步 H2D 拷贝,并等待前序 D2H event 完成,避免读写竞争。
  • release_device(sid):释放 device buffer;若 H2D 仍在飞行则同步等待。
  • release_host(sid):归还 host buffer 到 PinnedMemoryPool,等待相关 event 完成。
  • wait_for_transfers():让当前计算流等待拷贝流,用于 __exit__ 安全退出。
  • sync_all_transfers():异常路径下同步拷贝流,确保资源状态一致。

所有 public 方法以 storage_id 为参数,保持与逻辑 tensor / ShadowTensor 的解耦。

PhysicalBuffer

最小化物理状态机:

(device_buffer, device_event)  <──>  (host_buffer, host_event)
  • device_storage():返回 device resident storage;必要时从 host demand-page;等待 device_event 以确保 H2D 完成。
  • 事件同步遵循 "先等待对应 event,再访问 buffer" 原则,避免跨 stream 竞争。
PinnedMemoryPool

全局固定主机内存池:

  • 采用桶化 (bucket-based) 管理,桶大小为 2^10 ~ 2^31 bytes。
  • acquire(size):优先复用同桶空闲 buffer;无可用时按桶对齐申请新 pin_memory=True 内存;超过 max_host_bytes 则降级为普通 pageable 内存。
  • release(tensor, event):将 buffer 加入 pending 列表等待 event 完成后再回收,或立即放入可用池。
  • 线程安全:通过 threading.Lock 保护池状态。
BandwidthEstimator

profile_transfer_bandwidth() 在 Warmup 结束时通过 16 MiB 的 dummy buffer 实测 D2H / H2D 带宽,为 planner 提供真实硬件参数。若加速器不可用或测试失败,则回退默认值 16 Gbps。


5. 关键实现细节 (Key Implementation Details)

5.1 TorchDispatchMode 与生命周期

ActivationDispatchMode 继承自 torch.utils._python_dispatch.TorchDispatchMode,在 __torch_dispatch__ 中把每个算子转发给当前 executor 的 dispatch 方法。该模式在 OffloadSession.__enter__ 时启用,__exit__ 时退出,覆盖范围仅限于 session 上下文内的 eager 执行。

5.2 Opaque Region 的反向图连续性

@skip_offload 装饰的函数被包装为单个虚拟 op。由于函数输出可能是普通 Tensor,需要被转换为 ShadowTensor 并参与 autograd,我们在 OpaqueRegionEnd 这个 autograd.Function 内部完成 ShadowTensor 包裹。这解决了 wrapper subclass 在 autograd 图中的正确链接问题。

OpaqueRegionStart / OpaqueRegionEnd 的 backward 钩子负责:

  • 在反向进入 Opaque Region 时调用 enter_opaque_region(),避免内部算子触发额外的虚拟 op 记录。
  • 在反向退出 Opaque Region 时调用 exit_opaque_region() 并执行 on_op_end,完成反向虚拟 op 的 trace/ShadowTensor 处理。
5.3 内存安全与事件同步
  • D2H 与 H2D 的读写竞争copy_h2d 显式等待 buffer.host_event,确保 H2D 读取 host buffer 时,前序 D2H 已完成。
  • device buffer 回收安全release_device 在 H2D 仍在飞行时同步 device_event;D2H 飞行期间 device buffer 通过 record_stream(copy_stream) 防止被缓存分配器回收。
  • host buffer 回收安全release_host 将 buffer 与相关 event 一起放入 pending 队列,event 完成后才回收入可用池。
  • 异常安全__exit__ 在发生异常时调用 sync_all_transfers() + reset(),确保拷贝流与物理 buffer 状态被清理。
5.4 与现有 activation_checkpoint.swap 的关系
  • HyperOffload 作为独立包 hyper_parallel/auto_parallel/hyper_offload/ 引入,默认不启用。
  • swap API(MUST_SWAPSwapManagerSwapGroup)继续保留,用户可平滑迁移。
  • 推荐策略:在新模型/长序列训练中使用 OffloadSession 做全局自动卸载;在已有手工优化场景可继续用 swap 作为补充。

6. 测试与验证计划 (Test Plan)

6.1 单元测试 (Unit Tests)

tests/ut/auto_parallel/hyper_offload/ — 不依赖加速器,CPU 可运行

测试文件 覆盖内容
api/test_config.py OffloadConfig 构造、默认值、自定义参数
api/test_ir.py ActivationTraceResidencyScheduleOpGuideAccessKind 等 IR 数据结构
api/test_opaque.py @skip_offload 装饰器、Opaque Region 前后向、嵌套与空 session 行为;端到端 MLP/Transformer block 中验证 @skip_offload 行为与精度(含装饰器生成虚拟 op、replay 透传、同函数多次调用等)
api/test_session.py OffloadSession 生命周期、配置透传、warmup→replay 切换、异常清理
execution/test_base.py BaseExecutor 抽象行为、dispatch 流程、opaque op 包裹
execution/test_replay.py ReplayExecutor 按 schedule 执行动作、输出结构校验、action 类型异常
execution/test_tensor.py ShadowTensor 构造、resolve()、设备/host 回退、梯度传播
execution/test_tracker.py ActivationTracker storage 身份识别与生命周期追踪
execution/test_warmup.py WarmupExecutor 在线驱逐与 trace 记录
planning/test_greedy_planner.py GreedyResidencyPlanner 基本预算满足、write-before-copy 安全、retained sid 处理
runtime/test_bandwidth.py profile_transfer_bandwidth 带宽探测正确性
runtime/test_pinned_memory.py PinnedMemoryPool 桶化管理、申请/释放、线程安全
runtime/test_residency.py ResidencyManager D2H/H2D、release、跨 stream 事件同步、异常路径
runtime/test_timer.py DeviceTimer 计时器正确性
6.2 集成测试 (Integration Tests)

tests/torch/auto_parallel/hyper_offload/ — 需要 CUDA 或等效加速器

测试文件 覆盖内容
test_memory.py 在真实 CUDA 设备上设置严苛显存预算,验证 peak memory 低于基线且不触发 OOM
test_precision.py FP16/BF16/FP32 混合精度场景下,对比 Offload 与基线的 loss、梯度,要求严格一致
6.3 性能验证
  • 在长序列(如 8K/32K/128K)LLM 训练任务上,测量开启 HyperOffload 后的:
    • 峰值设备内存 (peak device memory)
    • 端到端 step time / throughput
    • PCIe 带宽利用率
  • 预期:在显存受限场景下可显著扩展可训练序列长度,传输开销被计算掩盖,吞吐损失 < 10%(具体取决于 PCIe 带宽与计算强度)。
6.4 兼容性验证
  • torch.compiletorch.autograd.Function、自定义算子、activation_checkpoint 的联合使用。
  • 多卡 DP/FSDP/TP 组合场景下的初步验证(当前版本主要面向单卡/数据并行 rank 的本地激活值)。

7. 接口变更与向后兼容 (API Compatibility)

本次变更仅新增接口,不修改现有接口:

# 新增公共 API
from hyper_parallel.auto_parallel.hyper_offload import OffloadConfig, OffloadSession, skip_offload
  • OffloadConfigOffloadSessionskip_offload 为新引入的公共符号。
  • 无现有函数签名变更、无行为回归。
  • activation_checkpoint.swap 保持原语义,用户可按需选择。

8. 已知限制与未来工作 (Limitations & Future Work)

  1. PyTorch-only(首版):MindSpore / 其他后端需要单独的 dispatch adapter。
  2. 静态图假设:当前 Planner 假设 Warmup trace 与后续 Replay 执行图结构完全一致。若存在 input-dependent 动态分支,需使用 @skip_offload 包裹。
  3. Planner 可扩展性:当前仅实现贪心 planner。后续可引入 ILP / 动态规划 / 机器学习 cost model,以在复杂 memory/compute 约束下获得更优调度。
  4. 多设备与分布式:当前 ResidencyManager 管理单个 rank 的 device/host 内存。跨 rank 协同卸载、与 FSDP/TP/PP 的深度融合是下一步方向。
  5. 参数/优化器卸载:本模块的 runtime 层可复用为 Parameter Offload / Optimizer Offload 的基础。

10. 参考文档 (References)


欢迎社区及 Maintainers 评审、提问与建议!

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

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 reviewing the proposed hyper_parallel/auto_parallel/hyper_offload/ API, IR, execution, planning, and runtime layers, then inspect tests/ut/auto_parallel/hyper_offload/. The work is complete when the RFC’s OffloadSession, tracing, planning, replay, residency management, and CPU-runnable validation plan are implemented with numerical equivalence and safe transfer cleanup.

Written by the indexing model from the issue text.

Assessment

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

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.