mindspore-ai / mindspore-ai/hyper-parallel
[RFC] HyperOffload:自动异步激活值卸载
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_checkpoint 的 swap 机制(通过 CheckpointPolicy.MUST_SWAP 与 SwapManager 手动包裹特定算子/模块),但该方案在工程实践中暴露出以下痛点:
- 侵入性高:用户需要深入理解网络结构,手动挑选并包裹待卸载的算子或子模块。模型迭代时,这些包裹点需要重新调优,维护成本随模型复杂度急剧上升。
- 粒度较粗:以模块(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" 进行卸载,使得传输次数与暴露的传输延迟最小化。
算法步骤:
- 将
ActivationTrace按 storage ID 分组,得到每个 storage 的访问序列(按 op_id 排序)。 - 计算每个 storage 在每个 op 上的驻留字节贡献,得到
resident_bytes[op_id]。 - 对每对相邻访问
(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则视为有效候选。
- 将候选按
(distance, size, -release_start.op_id)降序排序:优先卸载 "距离下次使用最远、尺寸最大" 的 activation。 - 依次选择候选,只要该候选覆盖的任一 op 当前
resident_bytes[op_id]仍超过预算,就将其加入调度表,并扣减相应 op 的驻留字节。 - 为每个被选中的候选生成:
COPY_D2H@copy_start.op_idRELEASE_DEVICE@release_start.op_idCOPY_H2D@end.op_id
- 对非 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^31bytes。 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/引入,默认不启用。- 旧
swapAPI(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)
- PyTorch-only(首版):MindSpore / 其他后端需要单独的 dispatch adapter。
- 静态图假设:当前 Planner 假设 Warmup trace 与后续 Replay 执行图结构完全一致。若存在 input-dependent 动态分支,需使用
@skip_offload包裹。 - Planner 可扩展性:当前仅实现贪心 planner。后续可引入 ILP / 动态规划 / 机器学习 cost model,以在复杂 memory/compute 约束下获得更优调度。
- 多设备与分布式:当前
ResidencyManager管理单个 rank 的 device/host 内存。跨 rank 协同卸载、与 FSDP/TP/PP 的深度融合是下一步方向。 - 参数/优化器卸载:本模块的 runtime 层可复用为 Parameter Offload / Optimizer Offload 的基础。
10. 参考文档 (References)
- PyTorch
TorchDispatchMode文档:https://pytorch.org/docs/stable/notes/extending.html - PyTorch Tensor Subclass (
__torch_dispatch__):https://pytorch.org/docs/stable/notes/extending.html#extending-torch-with-a-tensor-like-type - HyperParallel
activation_checkpoint.swap既有实现:platform/torch/activation_checkpoint/
欢迎社区及 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
- Read the whole issue, then the project's contributing guide.
- Comment on the issue to say you are picking it up — it saves two people doing the same work.
- Fork the repository and make your change on a branch.
- 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