Prhub

#44574 Preserve layout-changing clones

原始 PR 作者 mikekg 合并时间 2026-06-06 08:45 文件变更 2 提交数 1 评论 1 代码增减 +48 / -0

执行摘要

修复 unsafe clone 消除 pass 误删布局变更的 clone

修复 Mamba padded/sliced 输出场景下 .contiguous() 被错误移除的问题,该问题导致下游自定义算子因内存布局不兼容而行为异常。PR body 明确提到:"This fixes the Mamba padded/sliced output case where .contiguous() is required to produce the compact layout expected by downstream custom ops."

该 PR 虽然仅有 48 行变更,但修复了一个编译器优化 pass 的安全性问题,值得相关开发者精读,尤其是涉及自定义算子或依赖 tensor 布局的场景。

讨论亮点

该 PR 仅有一条来自 ProExpertProg 的评论("Thanks for the fix!"),无实质讨论或争议。两位 reviewer(ProExpertProgmgoin)均为 APPROVED,说明变更简单且得到认可。

实现拆解

  1. vllm/compilation/passes/ir/clone_elimination.py 中新增 clone_preserves_layout 函数:该函数从 fx.Node 的 meta['val'] 中提取 stride 和 storage offset,比较 clone 节点与原始节点的布局是否一致。若任一 meta 缺失或比较过程中抛出异常,保守返回 True(允许消除)。
  2. UnsafeCloneEliminationPass.__call__ 主循环中增加提前检查:在原有 clone 消除逻辑(写入检测、donated input 检查)之前,先调用 clone_preserves_layout;若布局已改变,则直接跳过该节点,保留 clone。
  3. tests/compile/passes/ir/test_clone_cleanup.py 中新增 test_keep_clone_that_changes_layout 测试用例:构造一个 x[:, :3].contiguous() 的 FX 图,确认 clone 消除 pass 执行后 clone 数量仍为 1,且输出 stride 为 (3, 1)
文件 模块 状态 重要度
vllm/compilation/passes/ir/clone_elimination.py 编译优化 modified 7.42
tests/compile/passes/ir/test_clone_cleanup.py 测试 modified 5.9

关键符号

clone_preserves_layout UnsafeCloneEliminationPass.__call__ test_keep_clone_that_changes_layout

关键源码片段

vllm/compilation/passes/ir/clone_elimination.py core-logic

新增 `clone_preserves_layout` 函数,修改 `UnsafeCloneEliminationPass` 主循环,是本 PR 的核心逻辑所在。

# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
import torch
from torch import fx# ... (imports omitted for brevity)def clone_preserves_layout(node: fx.Node, original_node: fx.Node) -> bool:
    """
    检查 clone 操作是否保持了原始张量的内存布局(stride 和 storage offset)。
    如果布局已改变(如因 .contiguous() 或切片后 clone 产生紧凑布局),
    则返回 False,以保留该 clone 避免下游算子出错。
    """
    node_val = node.meta.get("val")
    original_val = original_node.meta.get("val")
    # 若 meta 缺失,保守假设布局不变(允许消除)
    if node_val is None or original_val is None:
        return True
​
    try:
        node_stride = tuple(node_val.stride())
        original_stride = tuple(original_val.stride())
        node_storage_offset = node_val.storage_offset()
        original_storage_offset = original_val.storage_offset()
    except (AttributeError, RuntimeError):
        return True
​
    return (
        node_stride == original_stride
        and node_storage_offset == original_storage_offset
    )
​
​
class UnsafeCloneEliminationPass(VllmInductorPass):
    # ... (__init__ unchanged)
​
    def __call__(self, graph: fx.Graph) -> None:
        # ... (prelude: node_to_idx, donated_input_ids)
​
        for node in graph.nodes:
            if not is_func(node, torch.ops.aten.clone.default):
                continue
​
            original_node = node.args[0]
            assert isinstance(original_node, fx.Node)
​
            # --- 新增:检查 layout 是否改变,若改变则保留 clone ---
            if not clone_preserves_layout(node, original_node):
                logger.debug(
                    "Clone removal not possible, clone changes layout: "
                    "original_node=%s node=%s",
                    original_node,
                    node,
                )
                continue
​
            # ... ( 原有写入检测、donated input 检查逻辑,略 )
tests/compile/passes/ir/test_clone_cleanup.py test-coverage

新增测试用例 `test_keep_clone_that_changes_layout`,覆盖了布局变更 clone 的保护场景。

# ... (imports and test class)def test_keep_clone_that_changes_layout(self, clone_cleanup_pass):
    """
    Clone must be kept when it materializes a compact slice layout.
    本用例模拟 Mamba 场景中切片后调用 .contiguous() 产生紧凑布局的场景,
    验证 unsafe clone elimination pass 不会错误消除该 clone。
    """
    def f(x: torch.Tensor) -> torch.Tensor:
        return x[:, :3].contiguous()
​
    inp = torch.randn(4, 5)
    graph_module = make_fx(f)(inp)
    assert count_clones(graph_module.graph) == 1 # 确保初始图中有 1 个 clone
​
    expected = graph_module(inp)
    clone_cleanup_pass(graph_module.graph)
    graph_module.recompile()
    actual = graph_module(inp)
​
    # 核心断言:clone 仍被保留(未被错误消除)
    assert count_clones(graph_module.graph) == 1
    # 确认输出 stride 为紧凑布局 (3, 1) 而非视图的 (5, 1)
    assert actual.stride() == expected.stride() == (3, 1)
    torch.testing.assert_close(actual, expected)

评论区精华

没有提炼出高价值讨论线程

当前评论区没有形成足够清晰的争议点或结论,后续有更多讨论时会体现在这里。

风险与影响

风险很低。新增的 clone_preserves_layout 在 meta 缺失或异常时保守返回 True(允许消除),不会引入新的功能破坏。仅影响 torch.compile 编译路径,且仅会减少 bug 场景,不会影响已正确工作的案例。

影响范围窄,仅作用于 torch.compile 编译流程中的 clone elimination pass。修复了 Mamba 模型(如 Nemotron Super)在 torch.compile 下的正确性。对无 torch.compile 的场景无影响。

核心路径变更 测试覆盖增加

关联 Issue

未识别关联 Issue

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

完整报告

参与讨论