Prhub

#46276 [BugFix] weights processing peak memory reduction for nvfp4 MoE layers

原始 PR 作者 thisisjimmyfb 合并时间 2026-07-11 10:05 文件变更 2 提交数 3 评论 6 代码增减 +36 / -29

执行摘要

降低 NVFP4 MoE 权重重排的峰值内存占用

PR body 明确指出:在 16GB VRAM 上服务 nvidia/Qwen3.6-35B-A3B-NVFP4 等 MoE 模型时,reorder_w1w3_to_w3w1 返回两个新的替换张量导致 OOM。需要实现原地 chunk swap 来缓解内存峰值。

建议精读。该 PR 展示了一个经典的用空间换时间到时间换空间的优化思路,review 中对正确性的深入讨论(原地操作的副作用)值得所有开发者注意。chunk swap 的实现也值得作为类似操作的参考。

讨论亮点

Review 中 mgoin 建议简化实现为固定 chunk size,避免动态计算可用内存,并提供了参考代码。waynehacking8 在 RTX PRO 6000 (SM120) 上验证了所有 24 个测试用例通过,并指出新实现的 bit-exact 结果。同时 waynehacking8 发现 reorder_w1w3_to_w3w1 变为原地操作后,测试中的 w1_bf16 被意外突变,导致后续 BF16 参考计算的权重顺序错误,但测试通过了——因为输入 damping 使得 silu 近似线性,掩盖了错误。作者 thisisjimmyfb 随后在测试中使用了 w1_bf16.clone() 避免副作用,并添加了 assert。

实现拆解

  1. 重构 reorder_w1w3_to_w3w1 函数vllm/model_executor/layers/quantization/utils/flashinfer_fp4_moe.py):将原先的 split + cat + contiguous 操作替换为 chunk-wise 原地 swap。算法以 64MB 为 transient 上限,chunk size = min(half, 64MB / bytes_per_row),在循环中逐块交换 weight 和 scale 的前后两半。修改后函数保证输入必须连续,并保持输出连续。
  2. 调整测试文件tests/kernels/moe/test_flashinfer_b12x_moe.py):删除重复的 _reorder_gate_up_to_up_gate 辅助函数,改为直接调用生产代码的 reorder_w1w3_to_w3w1,确保测试覆盖生产路径;同时修复了测试中 FlashInferB12xExperts 初始化时缺失 _fc2_input_scale 导致 process_weights_after_loading 失败的问题。
  3. 修复 test_flashinfer_b12x_moe_relu2 测试:移除该测试调用 fused_moe 时错误的 inplace=False 参数(该参数已被上游移除),使测试通过。
文件 模块 状态 重要度
vllm/model_executor/layers/quantization/utils/flashinfer_fp4_moe.py 量化工具 modified 7.1
tests/kernels/moe/test_flashinfer_b12x_moe.py MoE 测试 modified 5.86

关键符号

reorder_w1w3_to_w3w1

关键源码片段

vllm/model_executor/layers/quantization/utils/flashinfer_fp4_moe.py core-logic

核心修改文件,将 `reorder_w1w3_to_w3w1` 从 O(2n) 空间复杂度优化为 O(n/2) 的原地 chunk swap。

def reorder_w1w3_to_w3w1(
    weight: torch.Tensor, scale: torch.Tensor, dim: int = -2
) -> tuple[torch.Tensor, torch.Tensor]:
    """Re-order concatenated `[w1, w3]` tensors to `[w3, w1]` in-place.    `weight` and `scale` must be contiguous; they remain contiguous on return.
    """
    assert weight.is_contiguous(), "weight must be contiguous"
    assert scale.is_contiguous(), "scale must be contiguous"
    size = weight.size(dim)
    assert size % 2 == 0, f"Expected even size in dim {dim}, got {size}"
    half = size // 2
    d = dim % weight.dim()
​
    # 以 64 MB 作为临时内存开销上限
    bytes_per_row = max(
        weight.numel() // size * weight.element_size(),
        scale.numel() // size * scale.element_size(),
    )
    chunk = max(1, min(half, (64 << 20) // max(bytes_per_row, 1)))
​
    # 构造索引切片
    fa, fb = [slice(None)] * weight.dim(), [slice(None)] * weight.dim()
    for off in range(0, half, chunk):
        end = min(off + chunk, half)
        fa[d], fb[d] = slice(off, end), slice(half + off, half + end)
        a, b = tuple(fa), tuple(fb)
        # 同时交换 weight 和 scale 的前后半部分
        for t in (weight, scale):
            tmp = t[b].clone() # 只复制一个 chunk
            t[b] = t[a]
            t[a] = tmp
​
    return weight, scale

评论区精华

实现简化:固定 chunk size 替代动态计算 设计

mgoin 建议使用固定 chunk size(64MB)替代动态计算可用内存,认为 overhead 可以忽略。

结论:作者采纳了建议,实现了基于 64MB 上限的固定 chunk size。 · 已解决

原地操作导致 w1_bf16 被意外突变 正确性

waynehacking8 发现测试中 `reorder_w1w3_to_w3w1` 是原地操作,导致后续 BF16 参考计算的 `w1_bf16` 被修改为 `[up, gate]` 顺序,但测试仍然通过(输入 damping 使 silu 近似线性,掩盖了错误)。

结论:作者在测试中改用 `w1_bf16.clone()` 传递副本,避免原地操作的副作用。 · 已解决

测试结果验证 测试

waynehacking8 在 RTX PRO 6000 (SM120) 上验证所有 24 个测试用例通过,且输出与旧版本 bit-exact。

结论:测试结果被接受,证明了实现正确。 · 已解决

风险与影响

风险较低。主要风险是原地操作可能意外变异调用方传入的张量(已在测试中用 clone 解决);64MB chunk size 是硬编码上限,对超大张量(>数百GB)可能产生较多迭代,但开销仍可控。没有破坏现有接口签名。

直接影响:reorder_w1w3_to_w3w1 调用者(主要在 FlashInfer NVFP4 权重加载路径)将获得更低的峰值内存,使得在 16GB VRAM 等资源受限系统上可以加载更大的 NVFP4 MoE 模型。间接影响:测试的修复和增强提高了 FlashInfer B12x MoE 路径的测试可靠性。

核心路径变更 原地修改语义

关联 Issue

未识别关联 Issue

当前没有检测到明确关联的 Issue 链接,后续同步到相关引用后会出现在这里。

完整报告

参与讨论