在 HyperParallel 框架中训练大规模 Transformer / LLM 时,前向传播产生的中间激活值 (activations) 是长序列 (long-context) 场景下设备内存 (Device Memory) 不足并触发 OOM 的核心原因之一。随着序列长度、批量大小和模型层数的同步增长,激活值占用的显存往往呈线性甚至超线性膨胀,严重制约了可训练模型的规模上限。
HyperParallel
项目中虽已存在基于 activation_checkpoint 的 swap 机制(通过 CheckpointPolicy.MUST_SWAP 与 SwapManager 手动包裹特定算子/模块),但该方案在工程实践中暴露出以下痛点:
activation_checkpoint
swap
CheckpointPolicy.MUST_SWAP
SwapManager
因此,我们需要一种零侵入、细粒度、自动异步卸载的新机制:HyperOffload。
HyperOffload
with OffloadSession(config):
TorchDispatchMode
@skip_offload
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 │ └─────────────────────────────────────────────────────┘
核心设计哲学:
GreedyResidencyPlanner
ResidencySchedule
ReplayExecutor
COPY_D2H
RELEASE_DEVICE
COPY_H2D
RELEASE_HOST
ResidencyManager
PinnedMemoryPool
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)
replay
OffloadSession 通过 contextvars.ContextVar 维护当前活跃 session,供 @skip_offload 查询。
contextvars.ContextVar
skip_offload
装饰器/透明 API,用于标记一段代码为 Opaque Region(虚拟 op)。被装饰函数内部的算子不再被逐个追踪,而是作为单个虚拟 op 记录。适用于动态控制流、第三方库函数或用户自定义算子。
@skip_offload def my_custom_block(x): # 内部大量细粒度算子不会被单独追踪 return some_dynamic_logic(x)
hyper_parallel/auto_parallel/hyper_offload/ir/
Warmup 阶段的完整记录,包含:
ops: list[TraceOp]
READ
WRITE
storage_sizes: dict[sid, bytes]
retained_sids: set[int]
ShadowTensor
memory_limit_bytes
d2h_bandwidth_gbps
h2d_bandwidth_gbps
Planner 输出,按 op 索引维护:
pre_actions[op_id]
post_actions[op_id]
ReplayExecutor 使用的预消化结构,避免在热路径上重复解析 ActivationTrace。每个 op 保存输出叶子数量与 leaf_index -> storage_id 的绑定关系,用于快速 ShadowTensor 包裹与结构校验。
leaf_index -> storage_id
hyper_parallel/auto_parallel/hyper_offload/execution/
BaseExecutor
定义统一的生命周期钩子:
on_op_begin(func, args, kwargs)
on_op_end(result)
dispatch(func, args, kwargs)
on_op_begin -> func -> on_op_end
execute_opaque_op(...)
OpaqueRegionStart
OpaqueRegionEnd
autograd.Function
on_op_begin
copy_d2h
release_device
on_op_end
ActivationTracker
func._schema.is_mutable
TraceOp
finish()
OpGuide.output_leaf_count
一个 torch.Tensor 子类(通过 _make_wrapper_subclass 构造),本身不持有 device 数据,而是持有对 PhysicalBuffer 的引用。每次被调度时通过 resolve() 从 PhysicalBuffer.device_storage() 重新构造视图:
torch.Tensor
_make_wrapper_subclass
PhysicalBuffer
resolve()
PhysicalBuffer.device_storage()
ShadowTensor 不缓存长期 device 引用,因此不会阻止底层 storage 被释放。
hyper_parallel/auto_parallel/hyper_offload/planning/
核心离线算法,目标是在满足 memory_limit_bytes 的前提下,选择一组 "access gap" 进行卸载,使得传输次数与暴露的传输延迟最小化。
算法步骤:
resident_bytes[op_id]
(release_start, end)
_EvictionCandidate
distance = end.op_id - release_start.op_id
copy_start
release_start
copy_start.op_id < end.op_id - 1
(distance, size, -release_start.op_id)
copy_start.op_id
release_start.op_id
end.op_id
复杂度:设 op 数为 S,storage 数为 N,候选数为 M。排序 O(M log M),模拟 O(M * S),对典型训练图足够高效;若后续需要,可引入线段树 / 差分数组优化到 O(M log S)。
S
N
M
O(M log M)
O(M * S)
O(M log S)
hyper_parallel/auto_parallel/hyper_offload/runtime/
物理驻留控制器,维护 storage_id -> PhysicalBuffer 映射:
storage_id -> PhysicalBuffer
bind(sid, tensor)
UntypedStorage
copy_d2h(sid)
copy_stream
copy_h2d(sid)
release_device(sid)
release_host(sid)
wait_for_transfers()
sync_all_transfers()
所有 public 方法以 storage_id 为参数,保持与逻辑 tensor / ShadowTensor 的解耦。
storage_id
最小化物理状态机:
(device_buffer, device_event) <──> (host_buffer, host_event)
device_storage()
device_event
全局固定主机内存池:
2^10 ~ 2^31
acquire(size)
pin_memory=True
max_host_bytes
release(tensor, event)
threading.Lock
BandwidthEstimator
profile_transfer_bandwidth() 在 Warmup 结束时通过 16 MiB 的 dummy buffer 实测 D2H / H2D 带宽,为 planner 提供真实硬件参数。若加速器不可用或测试失败,则回退默认值 16 Gbps。
profile_transfer_bandwidth()
ActivationDispatchMode 继承自 torch.utils._python_dispatch.TorchDispatchMode,在 __torch_dispatch__ 中把每个算子转发给当前 executor 的 dispatch 方法。该模式在 OffloadSession.__enter__ 时启用,__exit__ 时退出,覆盖范围仅限于 session 上下文内的 eager 执行。
ActivationDispatchMode
torch.utils._python_dispatch.TorchDispatchMode
__torch_dispatch__
dispatch
OffloadSession.__enter__
@skip_offload 装饰的函数被包装为单个虚拟 op。由于函数输出可能是普通 Tensor,需要被转换为 ShadowTensor 并参与 autograd,我们在 OpaqueRegionEnd 这个 autograd.Function 内部完成 ShadowTensor 包裹。这解决了 wrapper subclass 在 autograd 图中的正确链接问题。
OpaqueRegionStart / OpaqueRegionEnd 的 backward 钩子负责:
enter_opaque_region()
exit_opaque_region()
copy_h2d
buffer.host_event
record_stream(copy_stream)
release_host
reset()
activation_checkpoint.swap
hyper_parallel/auto_parallel/hyper_offload/
MUST_SWAP
SwapGroup
tests/ut/auto_parallel/hyper_offload/ — 不依赖加速器,CPU 可运行
tests/ut/auto_parallel/hyper_offload/
api/test_config.py
api/test_ir.py
AccessKind
api/test_opaque.py
api/test_session.py
execution/test_base.py
execution/test_replay.py
execution/test_tensor.py
execution/test_tracker.py
execution/test_warmup.py
planning/test_greedy_planner.py
runtime/test_bandwidth.py
profile_transfer_bandwidth
runtime/test_pinned_memory.py
runtime/test_residency.py
runtime/test_timer.py
DeviceTimer
tests/torch/auto_parallel/hyper_offload/ — 需要 CUDA 或等效加速器
tests/torch/auto_parallel/hyper_offload/
test_memory.py
test_precision.py
torch.compile
torch.autograd.Function
本次变更仅新增接口,不修改现有接口:
# 新增公共 API from hyper_parallel.auto_parallel.hyper_offload import OffloadConfig, OffloadSession, skip_offload
platform/torch/activation_checkpoint/
欢迎社区及 Maintainers 评审、提问与建议!
1. 背景与动机 (Background & Motivation)
在
HyperParallel框架中训练大规模 Transformer / LLM 时,前向传播产生的中间激活值 (activations) 是长序列 (long-context) 场景下设备内存 (Device Memory) 不足并触发 OOM 的核心原因之一。随着序列长度、批量大小和模型层数的同步增长,激活值占用的显存往往呈线性甚至超线性膨胀,严重制约了可训练模型的规模上限。项目中虽已存在基于
activation_checkpoint的swap机制(通过CheckpointPolicy.MUST_SWAP与SwapManager手动包裹特定算子/模块),但该方案在工程实践中暴露出以下痛点:因此,我们需要一种零侵入、细粒度、自动异步卸载的新机制:
HyperOffload。2. 目标与非目标 (Goals & Non-Goals)
2.1 目标 (Goals)
with OffloadSession(config):上下文即可开启激活值自动卸载。2.2 非目标 (Non-Goals)
HyperOffload聚焦中间激活值卸载,模型参数与优化器状态仍由 FSDP / HSDP / ZeRO 等既有模块管理。HyperOffload可与 checkpoint / recompute 叠加使用;它不是重计算,而是显存换带宽的互补技术。TorchDispatchMode精确追踪的分支/循环控制流,提供@skip_offload逃生舱,由用户显式标注。3. 设计概览 (Design Overview)
HyperOffload采用 "先追踪、后规划、再重放" 的两阶段执行范式:核心设计哲学:
TorchDispatchMode拦截所有 Tensor 操作,记录每个激活值 storage 的创建、读写、销毁时序。GreedyResidencyPlanner根据全局显存预算和 access interval,生成一个离线驻留调度表 (ResidencySchedule)。ReplayExecutor,严格按调度表执行COPY_D2H、RELEASE_DEVICE、COPY_H2D、RELEASE_HOST,并通过独立拷贝流与计算流同步实现重叠。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 # 可插拔规划器,默认 GreedyResidencyPlannerOffloadSession继承上下文管理器语义,内部维护两阶段 Executor:
__enter__时处于warmup模式,使用WarmupExecutor。__exit__时等待所有异步传输完成,调用WarmupExecutor.finish()得到ActivationTrace与OpGuide,再由planner.build(trace)生成ResidencySchedule,随后切换为ReplayExecutor。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/)ActivationTraceWarmup 阶段的完整记录,包含:
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 参考。ResidencySchedulePlanner 输出,按 op 索引维护:
pre_actions[op_id]:op 执行前需要完成的动作(目前主要是COPY_H2D预取)。post_actions[op_id]:op 执行后需要完成的动作(COPY_D2H、RELEASE_DEVICE、RELEASE_HOST)。OpGuideReplayExecutor使用的预消化结构,避免在热路径上重复解析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维持反向图连续性。WarmupExecutorActivationTrace与OpGuide。on_op_begin时执行 在线贪心驱逐:若当前驻留字节数超过memory_limit_bytes,则按 "最早产生 op 优先、同 op 大小大者优先" 的策略选择 victim,调用copy_d2h+release_device。on_op_end时:ActivationTracker识别本 op 新产生的 activation storage;func._schema.is_mutable)以标记 write access;TraceOp与OpGuide。finish()结束时进行带宽 profile,并返回 trace + guide。ReplayExecutorResidencySchedule执行 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()重新构造视图:ShadowTensor 不缓存长期 device 引用,因此不会阻止底层 storage 被释放。
4.4 规划层 (
hyper_parallel/auto_parallel/hyper_offload/planning/)GreedyResidencyPlanner核心离线算法,目标是在满足
memory_limit_bytes的前提下,选择一组 "access gap" 进行卸载,使得传输次数与暴露的传输延迟最小化。算法步骤:
ActivationTrace按 storage ID 分组,得到每个 storage 的访问序列(按 op_id 排序)。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。resident_bytes[op_id]仍超过预算,就将其加入调度表,并扣减相应 op 的驻留字节。COPY_D2H@copy_start.op_idRELEASE_DEVICE@release_start.op_idCOPY_H2D@end.op_idRELEASE_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_storage():返回 device resident storage;必要时从 host demand-page;等待device_event以确保 H2D 完成。PinnedMemoryPool全局固定主机内存池:
2^10 ~ 2^31bytes。acquire(size):优先复用同桶空闲 buffer;无可用时按桶对齐申请新pin_memory=True内存;超过max_host_bytes则降级为普通 pageable 内存。release(tensor, event):将 buffer 加入 pending 列表等待 event 完成后再回收,或立即放入可用池。threading.Lock保护池状态。BandwidthEstimatorprofile_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 钩子负责:enter_opaque_region(),避免内部算子触发额外的虚拟 op 记录。exit_opaque_region()并执行on_op_end,完成反向虚拟 op 的 trace/ShadowTensor 处理。5.3 内存安全与事件同步
copy_h2d显式等待buffer.host_event,确保 H2D 读取 host buffer 时,前序 D2H 已完成。release_device在 H2D 仍在飞行时同步device_event;D2H 飞行期间 device buffer 通过record_stream(copy_stream)防止被缓存分配器回收。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.pyOffloadConfig构造、默认值、自定义参数api/test_ir.pyActivationTrace、ResidencySchedule、OpGuide、AccessKind等 IR 数据结构api/test_opaque.py@skip_offload装饰器、Opaque Region 前后向、嵌套与空 session 行为;端到端 MLP/Transformer block 中验证@skip_offload行为与精度(含装饰器生成虚拟 op、replay 透传、同函数多次调用等)api/test_session.pyOffloadSession生命周期、配置透传、warmup→replay 切换、异常清理execution/test_base.pyBaseExecutor抽象行为、dispatch 流程、opaque op 包裹execution/test_replay.pyReplayExecutor按 schedule 执行动作、输出结构校验、action 类型异常execution/test_tensor.pyShadowTensor构造、resolve()、设备/host 回退、梯度传播execution/test_tracker.pyActivationTrackerstorage 身份识别与生命周期追踪execution/test_warmup.pyWarmupExecutor在线驱逐与 trace 记录planning/test_greedy_planner.pyGreedyResidencyPlanner基本预算满足、write-before-copy 安全、retained sid 处理runtime/test_bandwidth.pyprofile_transfer_bandwidth带宽探测正确性runtime/test_pinned_memory.pyPinnedMemoryPool桶化管理、申请/释放、线程安全runtime/test_residency.pyResidencyManagerD2H/H2D、release、跨 stream 事件同步、异常路径runtime/test_timer.pyDeviceTimer计时器正确性6.2 集成测试 (Integration Tests)
tests/torch/auto_parallel/hyper_offload/— 需要 CUDA 或等效加速器test_memory.pytest_precision.py6.3 性能验证
HyperOffload后的:6.4 兼容性验证
torch.compile、torch.autograd.Function、自定义算子、activation_checkpoint的联合使用。7. 接口变更与向后兼容 (API Compatibility)
本次变更仅新增接口,不修改现有接口:
# 新增公共 API from hyper_parallel.auto_parallel.hyper_offload import OffloadConfig, OffloadSession, skip_offloadOffloadConfig、OffloadSession、skip_offload为新引入的公共符号。activation_checkpoint.swap保持原语义,用户可按需选择。8. 已知限制与未来工作 (Limitations & Future Work)
@skip_offload包裹。ResidencyManager管理单个 rank 的 device/host 内存。跨 rank 协同卸载、与 FSDP/TP/PP 的深度融合是下一步方向。10. 参考文档 (References)
TorchDispatchMode文档:https://pytorch.org/docs/stable/notes/extending.html__torch_dispatch__):https://pytorch.org/docs/stable/notes/extending.html#extending-torch-with-a-tensor-like-typeactivation_checkpoint.swap既有实现:platform/torch/activation_checkpoint/欢迎社区及 Maintainers 评审、提问与建议!