XTuner Memory Optimization
显存峰值
一次训练 step 的峰值可以先用下面的对象账本理解:
$$
M_{\text{peak}} =
M_{\text{model state}} +
M_{\text{saved activations}} +
M_{\text{operator workspace}} +
M_{\text{communication buffers}} +
M_{\text{live inputs/outputs}} +
M_{\text{reserved unused}} +
M_{\text{runtime margin}}
$$
公式中的加法不是说这些对象一直同时存在,而是提醒我们:峰值取决于某一时刻有哪些对象重叠。同一个优化在不同 schedule、序列长度、并行组和互联带宽下,可能把峰值移到另一个阶段。
| 物理动作 | 回答的问题 | 本文技术 |
|---|---|---|
| Reduce | 真正减少需要计算或保存的数据量 | MTP 重计算、跳过 Pad Token |
| Bound | 把一次在途的大对象限制在块大小以内 | ChunkLoss、ChunkMoE |
| Shift | 把设备常驻对象搬到 CPU | ViT/LM 激活卸载、SwapAdamW、Muon 状态卸载、pixel values |
| Shorten lifetime | 让对象最后一次使用后尽快变成可回收 | MTP/LMHead 重排、roll tensor、Muon workspace、pixel values |
| Allocator cleanup | 归还已经空闲的缓存块 | gc + empty_cache |
分块减峰
ChunkLoss
语言模型输出层会把 [B, S, H] 投影为 [B, S, V]。当词表 V 很大时,logits、交叉熵 workspace 和局部梯度会在 loss 阶段形成尖峰。ChunkLoss 沿序列轴把 S 切成大小为 C 的块,每块完成投影、loss 和局部求导后再处理下一块。[^chunk-loss]
| 对象 | 如何变化 | 不会自动变化 |
|---|---|---|
| logits / CE workspace | 从约 O(B·S·V) 限制为 O(B·C·V) |
词表大小 V |
| 块输入与局部梯度 | 一块用完即可合并 | 完整 hidden 的总大小 |
| LMHead 权重梯度 | 分块求出后累加 | 模型参数与优化器状态 |
1 | def chunk_loss(hidden, labels, lm_head, chunk_size): |
适用范围与效果:
- 适合大词表、长序列且 profiler 显示 LMHead/loss 是峰值阶段的训练。
- 块越小不一定越好:会增加 kernel launch、Python/Autograd 调度与权重梯度累加次数。
- 效果是 Bound,不是全局按比例下降:模型状态、完整 hidden 和其他层激活仍构成下限。
ChunkMoE
MoE 的 dispatch、专家 MLP 和 combine 会产生 token 相关中间量。ChunkMoE 先按完整序列完成 Attention 与 Gate,再把 hidden、residual 和 router 结果切块,串行执行块级 MoE,最后拼接输出。[^chunk-moe]
| 对象 | 如何变化 | 仍然存在的下限 |
|---|---|---|
| 专家中间激活 | 一次只处理 S/K 的 token |
单块最大专家负载 |
| dispatch/combine workspace | 随块复用 | 通信元数据与路由结果 |
| Attention、Gate、输出 | 不因 MoE 切块而消失 | clone、结果列表和 concat |
1 | def chunk_moe_forward(hidden, mask, chunk_size): |
适用范围与效果:
- 适合峰值明确出现在 experts/dispatcher,且序列可安全切分的 MoE 路径。
- 依赖细粒度重计算边界:Attention 与 MoE 若被错误地一起 checkpoint,可能重复更多计算。
- 效果是 Bound:专家激活下降,但完整 Attention、Gate、输出、clone 与 concat 不能按块数等比例下降。
激活卸载
PyTorch Autograd 会保存反向所需的张量;saved_tensors_hooks 允许在保存时 pack、取用时 unpack。XTuner 的 offload manager 把选中的张量异步复制到 pinned CPU,随后把设备 storage 缩小为零;反向前再 H2D 恢复。[^saved-hooks][^activation-offload]
ViT 激活卸载
视觉 token 数会随图像分辨率、图片数和视频帧数快速增长。ViT offload 在视觉 block 范围内选择需要保存的输入激活,并使用 group="vision" 管理其键和预取顺序。[^vit-offload]
| 对象 | 设备侧变化 | 代价与边界 |
|---|---|---|
| ViT block saved input | D2H 后释放设备 storage | pinned CPU 容量增加 |
| offload event / key | 记录副本完成与恢复顺序 | 需要 stream/event 同步 |
| visual embeddings | 仍供 LM 使用 | 不属于被卸载的 block input |
1 | def vit_block_with_offload(hidden): |
适用范围与效果:
- 适合视觉激活占比高、CPU 内存充足且 PCIe/HCCS/NVLink 传输可与计算覆盖的 VLM。
- 效果是 Shift:设备 saved activations 下降,但总系统内存没有消失。
- 若反向很快追上前向或互联较慢,H2D 等待可能直接暴露在关键路径上。
LM 激活卸载
LM offload 与 ViT 使用同一物理机制,但对象变成 Decoder 层为反向保存的输入,分组为 text。层数越深、序列越长,前向末尾累积的 saved activations 越多。[^lm-offload]
| 对象 | 设备侧变化 | 仍需验证 |
|---|---|---|
| Decoder saved input | 前向后转移到 CPU | 是否只选中目标 hidden |
| CPU activation queue | 随层数累积 | 主机内存容量与 NUMA |
| H2D 恢复对象 | 反向到该层前短暂回到设备 | 预取是否覆盖传输 |
1 | def decoder_layer_with_offload(hidden, layer_id): |
适用范围与效果:
- 适合saved activations 是主要峰值且计算足够长,能够隐藏大部分传输的 Decoder。
- 不适合盲目全卸载:过多小张量会放大事件和副本管理成本。
- 实际收益由 CPU 容量、互联带宽、预取距离和计算覆盖共同决定。
优化器状态
SwapAdamW
AdamW 通常为每个参数保留一阶矩 m 和二阶矩 v;AMSGrad 还会增加 max_v。SwapAdamW 把 CPU pinned tensor 作为权威状态,step 时逐参数 H2D、更新,再 D2H 写回。源码虽然使用 non-blocking copy,但 step 末尾仍有显式同步,不能简单理解为“传输完全异步隐藏”。[^swap-adamw]
| 对象 | 常驻位置 | 更新窗口 |
|---|---|---|
exp_avg / exp_avg_sq |
pinned CPU | 当前参数 step 时 H2D |
max_exp_avg_sq |
AMSGrad 时 pinned CPU | 与其他状态一起搬运 |
| 参数与梯度 | 仍在其训练设备/分片位置 | AdamW kernel 消费 |
1 | def swap_adamw_step(params): |
适用范围与效果:
- 适合优化器状态是常驻显存大头、CPU 内存足够、step time 能接受传输的训练。
- 效果是 Shift:
m/v常驻显存下降,参数、梯度和当前参数状态仍需设备空间。 - 对小参数很多的模型,逐参数 copy 与 Python 循环可能比大块批处理更低效。
Muon 状态卸载
Muon 至少维护 momentum;XTuner 中部分参数组仍走 AdamW,并可能维护 variance。swap=True 时这些状态驻留 pinned CPU,按参数 shape/sharding 分批搬到设备;AsyncRuntime 最多允许 3 个更新任务同时在途。[^muon-offload]
| 对象 | 如何处理 | 峰值来源 |
|---|---|---|
| Muon momentum | pinned CPU 常驻 | 每批 H2D 临时副本 |
| AdamW variance | 对相应参数组卸载 | 与 momentum 共同在途 |
| 全参与 NS workspace | 设备上按批创建 | 最多 3 个异步任务叠加 |
1 | def muon_step(parameter_batches, max_inflight=3): |
适用范围与效果:
- 适合Muon/AdamW 状态常驻占用明显、通信与计算可覆盖状态传输的分布式训练。
- 效果是 Shift + Bound:常驻状态转到 CPU,但 bound 取决于批大小与并发任务数。
- 调小并发会降峰值但可能减少 overlap;调大并发可能把 offload 节省重新花在 workspace 上。
MTP 生命周期
MTP 重计算
Multi-Token Prediction 会增加额外 depth。若每个 MTP 层像普通层一样保存内部激活,前向末尾会叠加多个 depth 的 saved activations。activation checkpoint 只保留入口,在 backward 中重跑 forward,以计算换显存。[^checkpoint][^mtp-recompute]
| 对象 | 前向后是否保留 | 代价 |
|---|---|---|
| checkpoint 输入 | 保留 | 仍占入口张量空间 |
| MTP 内部中间量 | 丢弃 | backward 时重算 |
| RNG / 状态语义 | 需要一致 | 有状态算子需额外谨慎 |
1 | def mtp_forward(hidden, future_embeddings, mtp_layers): |
适用范围与效果:
- 适合MTP saved activations 占峰值,且额外 forward 计算可接受的训练。
- 要求函数可重复:依赖外部可变状态、随机数或设备状态的实现需要核对 checkpoint 语义。
- 397b 分支当前相关条件带有强制 checkpoint 的实现特例;迁移时应按目标框架重新设计开关,而不是照抄条件表达式。
MTP 与 LMHead 重排
这项优化不改变单个 tensor 的字节数,而是改变对象重叠:旧路径在每个 MTP depth 后立即计算 LMHead/loss;新路径先完成所有 MTP forwards,再在尾部集中计算各 depth loss,并调整 FSDP 预取链。源码能证明顺序变化,不能仅凭提交证明节省比例。[^mtp-reorder]
| 对象 | 旧顺序 | 新顺序 |
|---|---|---|
| MTP outputs | 逐 depth 产出并立刻消费 | 先形成列表 |
| LMHead/loss 临时量 | 与后续 MTP/FSDP 对象重叠 | 集中到尾部 |
| 单个对象大小 | 不变 | 不变 |
1 | def reordered_mtp_and_loss(hidden, mtp_layers, lm_head): |
适用范围与效果:
- 只适合原 schedule 确实让多个大对象形成不必要重叠的实现。
- 效果是 Shorten overlap:必须用 memory timeline 或 snapshot 验证峰值窗口是否真的被错开。
- 新顺序也可能延长
mtp_outputs列表的生命周期,因此不能只看“LMHead 被放到最后”这句话。
roll_packed_tensor 生命周期
Sequence Parallel 下,rolled_e[rank_slice] 是一个小 view,但仍共享完整 rolled_e 的底层 storage。对 rank-local slice 调用 clone() 后,它获得独立存储;再清空旧 _raw_inputs_embeds 引用,完整 storage 才可能回收。[^mtp-roll]
| 对象 | 优化前 | 优化后 |
|---|---|---|
完整 rolled_e storage |
被 rank-local view 间接引用 | view clone 后可释放 |
| rank-local tensor | view,无独立 storage | [B, S/P, H] 独立副本 |
| 旧 raw embeddings | context 继续引用 | 最后消费后清空 |
1 | def next_mtp_embeddings(seq_ctx, sp_rank, sp_size): |
适用范围与效果:
- 适合memory snapshot 显示小 view 保活大 storage 的生命周期问题。
- 效果是 Shorten lifetime:完整 storage 可更早回收,但新增一次局部分片 copy。
- 若下游本就需要完整 tensor,或 clone 后仍保留其他 view,收益会消失。
论文中的 MTP
论文证据和工程证据需要分开:DeepSeek-V3 报告说明 MTP 架构及其模型质量实验;XTuner 源码与提交说明具体对象如何保存、重算和释放。只有在固定模型、batch、并行配置和版本的 profiler 中,才能给出这些工程补丁的显存节省比例。[^deepseek-v3]
临时量与碎片
Muon workspace 释放
Muon 的 AGRS 路径会先后产生 ag_output、reshape view、full_params、Newton-Schulz 输出和 ReduceScatter 输入。若 Python 引用跨阶段保留,这些大 storage 会在同一任务内重叠;再乘上最多 3 个异步任务,就形成 optimizer step 峰值。补丁在每个对象最后一个消费者之后显式 del。[^muon-release]
| 对象 | 最后消费者 | 释放后的意义 |
|---|---|---|
ag_output / ag_reshaped |
全参重建 | storage 可被后续阶段复用 |
full_params |
Newton-Schulz | 不再与 NS 结果长期重叠 |
ns_results / ns_stacked |
ReduceScatter | 任务尾部及时退场 |
1 | def agrs_muon_update(local_params): |
适用范围与效果:
- 适合Python 引用让已经无用的大临时量跨阶段存活的路径。
- 效果是 Shorten lifetime:不减少 AllGather、NS 或 ReduceScatter 各自必需的 workspace。
- 如果 allocator 已能复用但峰值由 3 个任务并发决定,还需同时调小 batch 或
max_tasks。
gc 与 empty_cache
gc.collect() 只会清理已经不可达的 Python 对象;empty_cache() 只释放 caching allocator 中未占用的缓存块。PyTorch 官方文档明确说明,它不会增加 PyTorch 可用于活张量的显存,只可能在部分情况下缓解碎片。[^empty-cache]
原始清单给出的阈值逻辑来自 commit c5e1005,但该提交不是本文固定 e949653 的祖先;e949653 当前 trainer 对应位置只有周期性 gc.collect()。因此,“超过 33 GB 自动 empty_cache”应视为分支补丁建议,不能写成 397b 当前默认行为。[^allocator-branch][^trainer-gc]
| 对象 | gc.collect() |
empty_cache() |
|---|---|---|
| 仍被引用的 tensor | 不释放 | 不释放 |
| 不可达 Python 循环对象 | 可清理 | 不负责 |
| allocator 未占用缓存块 | 间接创造可回收条件 | 归还给驱动 |
1 | def maybe_clean_allocator(step, threshold_bytes): |
适用范围与效果:
- 适合
reserved >> allocated、snapshot 显示空闲块难以复用、且在阶段边界可承受同步的场景。 - 不是 OOM 根治方案:若 allocated 本身逼近容量,应找活张量、batch、序列或 workspace。
- 高频触发会丢失缓存复用收益,并可能增加同步和重新分配开销。
Token 与输入
MoE 跳过 Pad Token
Pad Token 仍会经过 router 并产生 top-k;若大量 Pad 集中到少数专家,会抬高局部专家 token 数、通信量和中间激活。XTuner 在 Attention 与 Gate 之后选出 nonpad_indices,只让有效 token 进入 MoE,再 scatter 回原序列并保留 Pad 位置的 hidden。[^skip-pad]
| 对象 | 从 T_total 变为 T_valid |
仍按完整序列 |
|---|---|---|
| dispatch token / expert input | 是 | 否 |
| expert activation / combine | 是 | 否 |
| Attention、router top-k、最终输出 | 否 | 是 |
1 | def moe_skip_pad(hidden, attention_mask): |
适用范围与效果:
- 适合padding ratio 高,且 Pad 路由导致专家负载或峰值异常的 MoE batch。
- 效果是 Reduce:MoE token 相关对象从
T_total缩到T_valid。 - 若数据已做高效 packing、几乎没有 Pad,索引与 scatter 成本可能抵消收益。
pixel_values 生命周期
多图或视频的 pixel_values 可能很大。关键不是 ViT 结束后调用一个通用释放函数,而是不要让 SequenceContext 提前把完整像素搬到设备并长期持有:原始像素留在 CPU,Vision 路径按 SP 需要切分后再 H2D;产生 visual/deepstack embeddings 后,LM 不再需要设备像素。[^pixel-values]
| 对象 | 位置与生命周期 | 不同于 |
|---|---|---|
原始 pixel_values |
CPU 常驻,Vision 使用窗口才 H2D | ViT saved activations |
| rank-local device pixels | SP 切分后短暂存在 | 完整 CPU 原始输入 |
| visual embeddings | ViT 后继续供 LM 使用 | 可立即释放的像素输入 |
1 | def multimodal_forward(seq_ctx): |
适用范围与效果:
- 适合多图、视频或高分辨率 VLM,且设备像素在 ViT 后仍被 context 引用的实现。
- 效果是 Shift + Shorten:像素输入不再提前、长期占设备,但 CPU 原始输入仍存在。
- ViT 激活卸载与 pixel lifetime 可以组合:前者处理反向保存对象,后者处理原始输入。
联合选择
不要把 13 个开关一次全开。先用对象账本和时间线找到峰值,再选择作用于该对象的最小组合:
| 观察到的峰值 | 优先验证 | 需要同时关注 |
|---|---|---|
| LMHead logits / CE 突刺 | ChunkLoss | chunk launch 与梯度累加 |
| experts / dispatcher 突刺 | ChunkMoE、skip Pad | Attention/Gate 下限、负载均衡 |
| 前向末尾 saved activations 高 | ViT/LM offload、MTP recompute | CPU 容量、传输、重算时间 |
| optimizer step 高 | 状态 offload、Muon workspace 释放 | 并发任务数、通信与 NS workspace |
| MTP 与 LMHead 交界高 | schedule reorder、roll lifetime | 列表/view 的真实生命周期 |
| VLM 开始后像素长期占用 | pixel values 延迟 H2D | visual embeddings 与 ViT 激活 |
reserved >> allocated |
gc + empty_cache | 先确认没有活张量泄漏 |
一个可靠的 A/B 顺序是:
- 固定语义:模型、global batch、micro-batch、序列、精度、并行组和数据完全相同。
- 记录对象:同时采集
allocated、reserved、max_memory_allocated、memory snapshot 和 profiler timeline。 - 一次改一类物理动作:先 Reduce/Bound,再评估 Shift,最后处理 lifetime 与 allocator。
- 检查峰值迁移:优化一个阶段后,新的峰值可能转移到通信、optimizer 或 backward。
- 同时验吞吐与精度:显存下降但 step time 大幅增加,或 loss/梯度语义变化,都不能算完成。
总结
XTuner 这组显存优化可以归结为三个问题:
- 谁太大:logits、专家激活、saved activations、优化器状态、像素输入或临时 workspace。
- 为什么同时活着:算法依赖、schedule 重叠、view 保活 storage、Python 引用,还是异步任务并发。
- 应该怎样退场:切块限制、跳过无效 token、checkpoint 重算、搬到 CPU、调整顺序、clone 解除共享,或在最后消费者后删除引用。
最重要的边界是:**empty_cache() 不能替代对象级优化,offload 也不等于减少总内存。** 只有把峰值拆回具体对象、存储位置和生命周期,才能判断某个补丁是否可迁移,以及它节省的是哪一部分显存。
[^chunk-loss]: XTuner chunk_loss.py,commit e949653。
[^chunk-moe]: XTuner moe_decoder_layer.py ChunkMoE,commit e949653;引入提交 2df5b5d。
[^saved-hooks]: PyTorch saved_tensors_hooks。
[^activation-offload]: XTuner activation_offload.py,commit e949653。
[^vit-offload]: XTuner Qwen3-VL modeling_vision.py,commit e949653。
[^lm-offload]: XTuner moe.py LM activation offload,commit e949653。
[^swap-adamw]: XTuner swap_adamw.py,commit e949653。
[^muon-offload]: XTuner muon.py 状态初始化与 swap,commit e949653;异步更新路径。
[^checkpoint]: PyTorch activation checkpointing。
[^mtp-recompute]: XTuner MTP recompute 提交 67bc238;e949653 的 MTP checkpoint 路径。
[^mtp-reorder]: XTuner MTP/LMHead 重排提交 65e9fcd。
[^mtp-roll]: XTuner MTP roll 生命周期提交 d5790e9;mtp/utils.py 当前实现。
[^deepseek-v3]: DeepSeek-V3 Technical Report。
[^muon-release]: XTuner Muon 临时量释放提交 17047f5。
[^empty-cache]: PyTorch torch.cuda.memory.empty_cache。
[^allocator-branch]: XTuner 分支补丁提交 c5e1005。
[^trainer-gc]: XTuner e949653 trainer 周期性 GC。
[^skip-pad]: XTuner 跳过 Pad Token 提交 2512534;e949653 当前实现。
[^pixel-values]: XTuner pixel values 生命周期提交 9573bab;SequenceContext 保留 CPU 像素;Vision 路径按需处理。