跳转至

XTuner Memory Optimization

导言

“降低显存”不是一种动作。它可能是在减少对象大小限制同时在途的对象数量把对象搬到 CPU缩短对象生命周期,也可能只是把 allocator 中未占用的缓存块归还给驱动。

本文固定到 XTuner 397b 分支 commit e949653,从一个第一次接触训练显存优化的读者视角,拆解原始清单中的 13 个技术点。每项都回答:大对象是什么、为什么形成峰值、执行时序怎样、数据流经过哪里、伪代码如何写、适用于什么条件,以及效果边界在哪里。

XTuner memory waterline intuition

认知插画:切块限制单次水位,offload 把对象搬到 CPU,checkpoint 用计算换保存空间,及时断开引用缩短对象寿命;缓存碎片则不是仍在使用的活张量。

显存峰值

一次训练 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

先分清 allocated 与 reserved

memory_allocated 更接近活张量实际占用,memory_reserved 还包含 caching allocator 留着复用的块。只有后者高而前者低时,碎片和缓存策略才更值得怀疑;二者都高时,应先找仍然存活的大张量。

分块减峰

ChunkLoss

语言模型输出层会把 [B, S, H] 投影为 [B, S, V]。当词表 V 很大时,logits、交叉熵 workspace 和局部梯度会在 loss 阶段形成尖峰。ChunkLoss 沿序列轴把 S 切成大小为 C 的块,每块完成投影、loss 和局部求导后再处理下一块。1

ChunkLoss five-view mechanism

ChunkLoss 五联图:物理前后对比、设计逻辑、执行流程、组件时序和张量数据流。
对象 如何变化 不会自动变化
logits / CE workspace 从约 O(B·S·V) 限制为 O(B·C·V) 词表大小 V
块输入与局部梯度 一块用完即可合并 完整 hidden 的总大小
LMHead 权重梯度 分块求出后累加 模型参数与优化器状态
def chunk_loss(hidden, labels, lm_head, chunk_size):
    grad_hidden = zeros_like(hidden)
    grad_weight = zeros_like(lm_head.weight)
    total_loss = 0

    for start in range(0, hidden.size(1), chunk_size):
        stop = min(start + chunk_size, hidden.size(1))
        h = hidden[:, start:stop].detach().requires_grad_(True)
        y = labels[:, start:stop]

        logits = lm_head(h)
        loss_i = cross_entropy(logits, y, reduction="sum")
        grad_h, grad_w = autograd.grad(
            loss_i, (h, lm_head.weight)
        )

        total_loss += loss_i.detach()
        grad_hidden[:, start:stop].copy_(grad_h)
        grad_weight.add_(grad_w)

    save_for_backward(grad_hidden, grad_weight)
    return total_loss

适用范围与效果:

  • 适合大词表、长序列且 profiler 显示 LMHead/loss 是峰值阶段的训练。
  • 块越小不一定越好:会增加 kernel launch、Python/Autograd 调度与权重梯度累加次数。
  • 效果是 Bound,不是全局按比例下降:模型状态、完整 hidden 和其他层激活仍构成下限。

ChunkMoE

MoE 的 dispatch、专家 MLP 和 combine 会产生 token 相关中间量。ChunkMoE 先按完整序列完成 Attention 与 Gate,再把 hidden、residual 和 router 结果切块,串行执行块级 MoE,最后拼接输出。2

ChunkMoE five-view mechanism

ChunkMoE 五联图:只有专家路径按块串行,Attention、Gate 和最终输出仍保留完整序列语义。
对象 如何变化 仍然存在的下限
专家中间激活 一次只处理 S/K 的 token 单块最大专家负载
dispatch/combine workspace 随块复用 通信元数据与路由结果
Attention、Gate、输出 不因 MoE 切块而消失 clone、结果列表和 concat
def chunk_moe_forward(hidden, mask, chunk_size):
    residual, attn_hidden = attention_forward(hidden, mask)
    topk_ids, topk_weights = gate(attn_hidden)
    outputs = []

    for slc in sequence_slices(attn_hidden, chunk_size):
        h_i = attn_hidden[slc].clone()
        r_i = residual[slc].clone()
        route_i = (topk_ids[slc].clone(), topk_weights[slc].clone())

        out_i = checkpoint_if_enabled(
            moe_dispatch_experts_combine,
            h_i,
            route_i,
        )
        outputs.append(out_i + r_i)

    return concat(outputs, dim=sequence_axis)

适用范围与效果:

  • 适合峰值明确出现在 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 恢复。34

ViT 激活卸载

视觉 token 数会随图像分辨率、图片数和视频帧数快速增长。ViT offload 在视觉 block 范围内选择需要保存的输入激活,并使用 group="vision" 管理其键和预取顺序。5

ViT activation offload five-view mechanism

ViT 激活卸载五联图:CPU 保存的是反向仍需要的视觉激活,不是删除数学依赖。
对象 设备侧变化 代价与边界
ViT block saved input D2H 后释放设备 storage pinned CPU 容量增加
offload event / key 记录副本完成与恢复顺序 需要 stream/event 同步
visual embeddings 仍供 LM 使用 不属于被卸载的 block input
def vit_block_with_offload(hidden):
    def pack(saved):
        if saved.data_ptr() != hidden.data_ptr():
            return saved
        cpu = empty_pinned_like(saved)
        d2h_stream.copy_(cpu, saved, non_blocking=True)
        key = register("vision", cpu, d2h_event())
        saved.untyped_storage().resize_(0)
        return SwapTensor(key, saved.shape, saved.dtype, saved.device)

    def unpack(swapped):
        if not isinstance(swapped, SwapTensor):
            return swapped
        return prefetch_h2d_and_wait(swapped.key)

    with saved_tensors_hooks(pack, unpack):
        return vit_block(hidden)

适用范围与效果:

  • 适合视觉激活占比高、CPU 内存充足且 PCIe/HCCS/NVLink 传输可与计算覆盖的 VLM。
  • 效果是 Shift:设备 saved activations 下降,但总系统内存没有消失。
  • 若反向很快追上前向或互联较慢,H2D 等待可能直接暴露在关键路径上。

LM 激活卸载

LM offload 与 ViT 使用同一物理机制,但对象变成 Decoder 层为反向保存的输入,分组为 text。层数越深、序列越长,前向末尾累积的 saved activations 越多。6

LM activation offload five-view mechanism

LM 激活卸载五联图:前向按层 D2H,反向按相反层序 H2D 预取。
对象 设备侧变化 仍需验证
Decoder saved input 前向后转移到 CPU 是否只选中目标 hidden
CPU activation queue 随层数累积 主机内存容量与 NUMA
H2D 恢复对象 反向到该层前短暂回到设备 预取是否覆盖传输
def decoder_layer_with_offload(hidden, layer_id):
    key = f"text_block_{layer_id}"

    def pack(saved):
        if saved.data_ptr() == hidden.data_ptr():
            cpu = async_copy_to_pinned_cpu(saved)
            register(key, cpu, current_d2h_event())
            saved.untyped_storage().resize_(0)
            return SwapTensor(key, saved.meta)
        return saved

    def unpack(value):
        if isinstance(value, SwapTensor):
            prefetch_previous_layer_if_possible(value.key)
            return restore_to_device_and_wait(value.key)
        return value

    with saved_tensors_hooks(pack, unpack):
        return decoder_layer(hidden)

适用范围与效果:

  • 适合saved activations 是主要峰值且计算足够长,能够隐藏大部分传输的 Decoder。
  • 不适合盲目全卸载:过多小张量会放大事件和副本管理成本。
  • 实际收益由 CPU 容量、互联带宽、预取距离和计算覆盖共同决定。

优化器状态

SwapAdamW

AdamW 通常为每个参数保留一阶矩 m 和二阶矩 v;AMSGrad 还会增加 max_v。SwapAdamW 把 CPU pinned tensor 作为权威状态,step 时逐参数 H2D、更新,再 D2H 写回。源码虽然使用 non-blocking copy,但 step 末尾仍有显式同步,不能简单理解为“传输完全异步隐藏”。7

SwapAdamW five-view mechanism

SwapAdamW 五联图:CPU 常驻优化器状态,设备只在当前参数更新窗口持有临时副本。
对象 常驻位置 更新窗口
exp_avg / exp_avg_sq pinned CPU 当前参数 step 时 H2D
max_exp_avg_sq AMSGrad 时 pinned CPU 与其他状态一起搬运
参数与梯度 仍在其训练设备/分片位置 AdamW kernel 消费
def swap_adamw_step(params):
    for p in params:
        if p.grad is None:
            continue

        cpu_m, cpu_v, cpu_max_v = cpu_state[p]
        dev_m = cpu_m.to(p.device, non_blocking=True)
        dev_v = cpu_v.to(p.device, non_blocking=True)
        dev_max_v = maybe_to_device(cpu_max_v, p.device)

        adamw_update(
            p, p.grad, dev_m, dev_v,
            max_exp_avg_sq=dev_max_v,
        )

        cpu_m.copy_(dev_m, non_blocking=True)
        cpu_v.copy_(dev_v, non_blocking=True)
        maybe_copy_back(cpu_max_v, dev_max_v)

    synchronize_device()

适用范围与效果:

  • 适合优化器状态是常驻显存大头、CPU 内存足够、step time 能接受传输的训练。
  • 效果是 Shiftm/v 常驻显存下降,参数、梯度和当前参数状态仍需设备空间。
  • 对小参数很多的模型,逐参数 copy 与 Python 循环可能比大块批处理更低效。

Muon 状态卸载

Muon 至少维护 momentum;XTuner 中部分参数组仍走 AdamW,并可能维护 variance。swap=True 时这些状态驻留 pinned CPU,按参数 shape/sharding 分批搬到设备;AsyncRuntime 最多允许 3 个更新任务同时在途。8

Muon state offload five-view mechanism

Muon 状态卸载五联图:常驻状态被搬走,但并发批次、通信、全参重建和 Newton-Schulz workspace 仍会抬高瞬时峰值。
对象 如何处理 峰值来源
Muon momentum pinned CPU 常驻 每批 H2D 临时副本
AdamW variance 对相应参数组卸载 与 momentum 共同在途
全参与 NS workspace 设备上按批创建 最多 3 个异步任务叠加
def muon_step(parameter_batches, max_inflight=3):
    runtime = AsyncRuntime(max_tasks=max_inflight)

    for batch in parameter_batches:
        runtime.submit(lambda batch=batch: update_batch(batch))

    runtime.wait_all()

def update_batch(batch):
    state = h2d(cpu_state[batch], non_blocking=True)
    full_params = communication_reconstruct(batch.params)

    if batch.algorithm == "muon":
        updated = muon_newton_schulz(full_params, state.momentum)
    else:
        updated = adamw_update(full_params, state.momentum, state.variance)

    scatter_updated_parameters(updated)
    async_d2h(cpu_state[batch], state)

适用范围与效果:

  • 适合Muon/AdamW 状态常驻占用明显、通信与计算可覆盖状态传输的分布式训练。
  • 效果是 Shift + Bound:常驻状态转到 CPU,但 bound 取决于批大小与并发任务数。
  • 调小并发会降峰值但可能减少 overlap;调大并发可能把 offload 节省重新花在 workspace 上。

MTP 生命周期

MTP 重计算

Multi-Token Prediction 会增加额外 depth。若每个 MTP 层像普通层一样保存内部激活,前向末尾会叠加多个 depth 的 saved activations。activation checkpoint 只保留入口,在 backward 中重跑 forward,以计算换显存。910

MTP recompute five-view mechanism

MTP 重计算五联图:减少的是 saved activations,不是 MTP 层参数或当次执行的 workspace。
对象 前向后是否保留 代价
checkpoint 输入 保留 仍占入口张量空间
MTP 内部中间量 丢弃 backward 时重算
RNG / 状态语义 需要一致 有状态算子需额外谨慎
def mtp_forward(hidden, future_embeddings, mtp_layers):
    outputs = []

    for layer in mtp_layers:
        hidden = checkpoint(
            layer,
            hidden,
            future_embeddings,
            use_reentrant=False,
            preserve_rng_state=True,
        )
        outputs.append(hidden)

    return outputs

适用范围与效果:

  • 适合MTP saved activations 占峰值,且额外 forward 计算可接受的训练。
  • 要求函数可重复:依赖外部可变状态、随机数或设备状态的实现需要核对 checkpoint 语义。
  • 397b 分支当前相关条件带有强制 checkpoint 的实现特例;迁移时应按目标框架重新设计开关,而不是照抄条件表达式。

MTP 与 LMHead 重排

这项优化不改变单个 tensor 的字节数,而是改变对象重叠:旧路径在每个 MTP depth 后立即计算 LMHead/loss;新路径先完成所有 MTP forwards,再在尾部集中计算各 depth loss,并调整 FSDP 预取链。源码能证明顺序变化,不能仅凭提交证明节省比例11

MTP and LMHead reorder five-view mechanism

MTP/LMHead 重排五联图:目标是错开 LMHead 梯度缓冲与其他阶段峰值,不是压缩 LMHead 本身。
对象 旧顺序 新顺序
MTP outputs 逐 depth 产出并立刻消费 先形成列表
LMHead/loss 临时量 与后续 MTP/FSDP 对象重叠 集中到尾部
单个对象大小 不变 不变
def reordered_mtp_and_loss(hidden, mtp_layers, lm_head):
    mtp_outputs = []

    # Phase 1: finish MTP forwards and use the intended FSDP prefetch chain.
    for depth, layer in enumerate(mtp_layers):
        prefetch_next_module(depth)
        hidden = checkpoint(layer, hidden)
        mtp_outputs.append(hidden)

    # Phase 2: calculate all depth losses at the tail.
    mtp_loss = 0
    for depth, depth_hidden in enumerate(mtp_outputs):
        mtp_loss += chunk_loss(
            depth_hidden,
            shifted_label(depth),
            lm_head,
            chunk_size,
        )

    return mtp_loss / len(mtp_outputs)

适用范围与效果:

  • 只适合原 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 才可能回收。12

MTP rolled tensor lifetime five-view mechanism

roll_packed_tensor 五联图:小 view 也可能保活完整底层 storage;clone 的目的不是省掉局部分片,而是解除共享。
对象 优化前 优化后
完整 rolled_e storage 被 rank-local view 间接引用 view clone 后可释放
rank-local tensor view,无独立 storage [B, S/P, H] 独立副本
旧 raw embeddings context 继续引用 最后消费后清空
def next_mtp_embeddings(seq_ctx, sp_rank, sp_size):
    rolled = roll_packed_tensor(
        seq_ctx._raw_inputs_embeds,
        seq_ctx.cu_seq_lens,
    )

    rank_slice = split_for_sequence_parallel(
        rolled, sp_rank, sp_size
    )
    local = rank_slice.clone()  # break the full-storage alias

    seq_ctx._raw_inputs_embeds = None
    del rank_slice, rolled
    return local

适用范围与效果:

  • 适合memory snapshot 显示小 view 保活大 storage 的生命周期问题。
  • 效果是 Shorten lifetime:完整 storage 可更早回收,但新增一次局部分片 copy。
  • 若下游本就需要完整 tensor,或 clone 后仍保留其他 view,收益会消失。

论文中的 MTP

DeepSeek-V3 MTP Figure 3

DeepSeek-V3 Technical Report Figure 3:MTP 通过顺序预测额外 token,把相邻 depth 的表示与未来 token embedding 组合。该图解释 MTP 数据依赖,不证明 XTuner checkpoint 或调度重排的显存收益。

DeepSeek-V3 MTP Table 4

DeepSeek-V3 Technical Report Table 4:论文设置下的 MTP 模型质量消融。它支持“为什么使用 MTP”的讨论,但不是 XTuner 重计算、LMHead 重排或 roll tensor 优化的显存 benchmark。

论文证据和工程证据需要分开:DeepSeek-V3 报告说明 MTP 架构及其模型质量实验;XTuner 源码与提交说明具体对象如何保存、重算和释放。只有在固定模型、batch、并行配置和版本的 profiler 中,才能给出这些工程补丁的显存节省比例。13

临时量与碎片

Muon workspace 释放

Muon 的 AGRS 路径会先后产生 ag_output、reshape view、full_params、Newton-Schulz 输出和 ReduceScatter 输入。若 Python 引用跨阶段保留,这些大 storage 会在同一任务内重叠;再乘上最多 3 个异步任务,就形成 optimizer step 峰值。补丁在每个对象最后一个消费者之后显式 del14

Muon workspace release five-view mechanism

Muon workspace 释放五联图:缩短同一 update batch 内的临时量重叠,但不改变各阶段自身大小和任务并发数。
对象 最后消费者 释放后的意义
ag_output / ag_reshaped 全参重建 storage 可被后续阶段复用
full_params Newton-Schulz 不再与 NS 结果长期重叠
ns_results / ns_stacked ReduceScatter 任务尾部及时退场
def agrs_muon_update(local_params):
    ag_output = all_gather(local_params)
    ag_reshaped = reshape_for_optimizer(ag_output)
    full_params = flatten_params(ag_reshaped)
    del ag_output, ag_reshaped

    ns_results = newton_schulz(full_params)
    del full_params

    ns_stacked = stack(ns_results)
    del ns_results

    local_update = reduce_scatter(ns_stacked)
    del ns_stacked
    return local_update

适用范围与效果:

  • 适合Python 引用让已经无用的大临时量跨阶段存活的路径。
  • 效果是 Shorten lifetime:不减少 AllGather、NS 或 ReduceScatter 各自必需的 workspace。
  • 如果 allocator 已能复用但峰值由 3 个任务并发决定,还需同时调小 batch 或 max_tasks

gc 与 empty_cache

gc.collect() 只会清理已经不可达的 Python 对象;empty_cache() 只释放 caching allocator 中未占用的缓存块。PyTorch 官方文档明确说明,它不会增加 PyTorch 可用于活张量的显存,只可能在部分情况下缓解碎片。15

原始清单给出的阈值逻辑来自 commit c5e1005,但该提交不是本文固定 e949653 的祖先;e949653 当前 trainer 对应位置只有周期性 gc.collect()。因此,“超过 33 GB 自动 empty_cache”应视为分支补丁建议,不能写成 397b 当前默认行为。1617

Allocator cleanup five-view mechanism

gc + empty_cache 五联图:活张量原样保留,只有引用已经消失并进入缓存的空闲块才可能归还驱动。
对象 gc.collect() empty_cache()
仍被引用的 tensor 不释放 不释放
不可达 Python 循环对象 可清理 不负责
allocator 未占用缓存块 间接创造可回收条件 归还给驱动
def maybe_clean_allocator(step, threshold_bytes):
    peak = device.max_memory_allocated()

    if step % cleanup_interval == 0 or peak > threshold_bytes:
        gc.collect()
        device.empty_cache()

    # This does not free any tensor that is still referenced.

适用范围与效果:

  • 适合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。18

MoE skip padding five-view mechanism

跳过 Pad Token 五联图:减少的是 MoE token 相关计算与通信,Attention 和 Gate 仍按完整序列执行。
对象 T_total 变为 T_valid 仍按完整序列
dispatch token / expert input
expert activation / combine
Attention、router top-k、最终输出
def moe_skip_pad(hidden, attention_mask):
    residual, hidden = attention_forward(hidden, attention_mask)
    topk_ids, topk_weights = gate(hidden)

    nonpad = where(attention_mask.reshape(-1) != 0)
    pad = where(attention_mask.reshape(-1) == 0)

    valid_hidden = index_select(hidden.reshape(-1, H), nonpad)
    valid_route = (
        index_select(topk_ids, nonpad),
        index_select(topk_weights, nonpad),
    )
    valid_out = dispatch_experts_combine(valid_hidden, valid_route)

    output = residual.reshape(-1, H).clone()
    output[nonpad] += valid_out
    output[pad] += hidden.reshape(-1, H)[pad]
    return output.reshape_as(hidden)

适用范围与效果:

  • 适合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 不再需要设备像素。19

pixel values lifetime five-view mechanism

pixel_values 生命周期五联图:延迟 H2D 并缩短设备驻留窗口,ViT 激活和 visual embeddings 则是另外两类对象。
对象 位置与生命周期 不同于
原始 pixel_values CPU 常驻,Vision 使用窗口才 H2D ViT saved activations
rank-local device pixels SP 切分后短暂存在 完整 CPU 原始输入
visual embeddings ViT 后继续供 LM 使用 可立即释放的像素输入
def multimodal_forward(seq_ctx):
    # Generic context-to-device logic intentionally skips pixel_values.
    move_metadata_and_text_tensors_to_device(seq_ctx)
    cpu_pixels = seq_ctx.pixel_values

    local_pixels = split_pixels_for_sp(cpu_pixels, sp_rank, sp_size)
    device_pixels = local_pixels.to(device, non_blocking=True)
    visual_embeds, deepstack_embeds = vision_encoder(device_pixels)

    del device_pixels, local_pixels
    return language_model(
        seq_ctx.text_inputs,
        visual_embeds,
        deepstack_embeds,
    )

适用范围与效果:

  • 适合多图、视频或高分辨率 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 顺序是:

  1. 固定语义:模型、global batch、micro-batch、序列、精度、并行组和数据完全相同。
  2. 记录对象:同时采集 allocatedreservedmax_memory_allocated、memory snapshot 和 profiler timeline。
  3. 一次改一类物理动作:先 Reduce/Bound,再评估 Shift,最后处理 lifetime 与 allocator。
  4. 检查峰值迁移:优化一个阶段后,新的峰值可能转移到通信、optimizer 或 backward。
  5. 同时验吞吐与精度:显存下降但 step time 大幅增加,或 loss/梯度语义变化,都不能算完成。

最小效果说明

每个实验至少记录“优化前峰值、优化后峰值、峰值所在阶段、step time、CPU 内存、传输时间、是否改变 batch/shape”。没有固定这些条件时,只能说明机制,不能声称节省了某个百分比。

总结

XTuner 这组显存优化可以归结为三个问题:

  1. 谁太大:logits、专家激活、saved activations、优化器状态、像素输入或临时 workspace。
  2. 为什么同时活着:算法依赖、schedule 重叠、view 保活 storage、Python 引用,还是异步任务并发。
  3. 应该怎样退场:切块限制、跳过无效 token、checkpoint 重算、搬到 CPU、调整顺序、clone 解除共享,或在最后消费者后删除引用。

最重要的边界是:empty_cache() 不能替代对象级优化,offload 也不等于减少总内存。 只有把峰值拆回具体对象、存储位置和生命周期,才能判断某个补丁是否可迁移,以及它节省的是哪一部分显存。

评论