已开启
[RFC] HyperOffload:自动异步激活值卸载 #211
xmoqian创建于  6月13日
xmoqian
6月13日 创建

1. 背景与动机 (Background & Motivation)

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

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

  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_D2H、RELEASE_DEVICE、COPY_H2D、RELEASE_HOST,并通过独立拷贝流与计算流同步实现重叠。
  • Residency:ResidencyManager 与 PinnedMemoryPool 负责物理 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() 得到 ActivationTrace 与 OpGuide,再由 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_D2H、RELEASE_DEVICE、RELEASE_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

  • 在线记录 ActivationTrace 与 OpGuide。
  • 在 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;
    • 生成 TraceOp 与 OpGuide。
  • finish() 结束时进行带宽 profile,并返回 trace + guide。

ReplayExecutor

  • 进入 replay 后,严格按 ResidencySchedule 执行 pre/post actions。
  • 在 on_op_begin 执行 COPY_H2D 预取。
  • 在 on_op_end 执行 COPY_D2H、RELEASE_DEVICE、RELEASE_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_start 为 release_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_SWAP、SwapManager、SwapGroup)继续保留,用户可平滑迁移。
  • 推荐策略:在新模型/长序列训练中使用 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 ActivationTrace、ResidencySchedule、OpGuide、AccessKind 等 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.compile、torch.autograd.Function、自定义算子、activation_checkpoint 的联合使用。
  • 多卡 DP/FSDP/TP 组合场景下的初步验证(当前版本主要面向单卡/数据并行 rank 的本地激活值)。

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

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

# 新增公共 API
from hyper_parallel.auto_parallel.hyper_offload import OffloadConfig, OffloadSession, skip_offload
  • OffloadConfig、OffloadSession、skip_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 评审、提问与建议!

likedislike
Xxmoqian
6月13日 修改了issue 的描述
Xxmoqian
6月13日 修改了issue 的描述
Xxmoqian
6月15日 修改了issue 的描述
Xxmoqian
6月16日 修改了issue 的描述
Xxmoqian
8月25日 关联了pull request:fix: defer host buffer reuse on in-flight H2D in clear_runtime