已开启
[Bug]: 【26.2.0众测】PyTorch 2.9.0环境下Guard Filter过滤运行时状态后仍触发重编译 #4856
Aimermo创建于  17 天前
Aimermo
17 天前 创建

在提交新问题之前,请确保您已经在社区中搜索过相关问题,并使用了社区中提供的资源/工具后,仍未找到满意的解决方式。

⚠️ 安全信息提醒:请仔细检查提供的文本内容,确保其不包含敏感数据信息,包括但不限于:

  • API 令牌或密钥
  • 密码或身份验证凭证
  • 私有网址或接口地址
  • 个人或机密数据
  • ...

在分享配置信息或代码示例时,请将敏感信息脱敏处理,或使用 <TOKEN> 等占位符替代原有内容。

环境信息

OS: Ubuntu 22.04, aarch64
Kernel: 5.10.0-216.0.0.115.oe2203sp4.aarch64
Python: 3.12.13
PyTorch: 2.9.0+cpu
torch_npu: 2.9.0.post6
CANN: 9.1.0
NPU: Ascend 910
Backend: AOT_Eager

🐛 问题描述

问题描述

通过torch.compile的guard_filter_fn过滤运行时状态Guard后,首次在torch.no_grad()下执行成功;切换到torch.enable_grad()并使用torch.compiler.set_stance("fail_on_recompile")再次执行时,仍然检测到重编译。

guard_filter_fn参数本身能被当前版本接受,问题集中在运行时状态Guard过滤未能阻止grad mode变化触发重编译。

完整最小复现代码

将以下内容保存为repro_guard_runtime.py:

import torch
import torch_npu

RUNTIME_GUARDS = {
    "GRAD_MODE",
    "TORCH_FUNCTION_STATE",
    "GLOBAL_STATE",
    "DEFAULT_DEVICE",
    "DETERMINISTIC_ALGORITHMS",
    "AUTOCAST_STATE",
    "FSDP_TRAINING_STATE",
}


def filter_runtime_guards(entries):
    return [
        entry.guard_type not in RUNTIME_GUARDS
        for entry in entries
    ]


class Model(torch.nn.Module):
    def __init__(self):
        super().__init__()
        self.linear = torch.nn.Linear(8, 8)

    def forward(self, x):
        return torch.relu(self.linear(x))


print("torch:", torch.__version__)
print("torch_npu:", torch_npu.__version__)
print("NPU available:", torch.npu.is_available())

model = Model().eval().to("npu:0")
x = torch.randn(2, 8, device="npu:0")

compiled = torch.compile(
    model,
    backend="aot_eager",
    options={"guard_filter_fn": filter_runtime_guards},
)

with torch.no_grad():
    first_output = compiled(x)
    torch.npu.synchronize()
print("first run: PASS")

try:
    with torch.compiler.set_stance("fail_on_recompile"):
        with torch.enable_grad():
            second_output = compiled(x)
            torch.npu.synchronize()
    print("recompile eliminated: PASS")
except RuntimeError as error:
    print("recompile eliminated: FAIL")
    print(type(error).__name__ + ":", error)

复现命令

先加载本机CANN环境,再运行脚本:

source /usr/local/Ascend/cann-9.1.0/set_env.sh
python repro_guard_runtime.py

如果CANN安装在其他位置,应将set_env.sh替换为实际路径;复现代码本身不依赖任何项目目录。

实际输出

torch: 2.9.0+cpu
torch_npu: 2.9.0.post6
NPU available: True
first run: PASS
recompile eliminated: FAIL
RuntimeError: Detected recompile when torch.compile stance is 'fail_on_recompile'.
filename: 'repro_guard_runtime.py', function name: 'forward'

上述最小代码已在本环境独立执行并稳定复现。

预期行为

GRAD_MODE等运行时状态Guard已被guard_filter_fn过滤后,从torch.no_grad()切换到torch.enable_grad()再次执行不应触发重编译,fail_on_recompile不应抛出异常。

对照现象

在同一环境中,按变量名过滤、过滤全局变量以及内置skip_guard_on_inbuilt_nn_modules_unsafe示例可以正常执行并消除对应重编译,说明guard_filter_fn入口可用。

欢迎加入社区,感谢您对社区的贡献 🎉!

likedislike
TorchNPU-BotTorchNPU-Bot成员
17 天前 添加了label:triage-review
TorchNPU-Bot
TorchNPU-Bot成员
17 天前 评论:

issue待分派,添加triage-review标签

likedislike
ascend-robotascend-robot成员
17 天前 添加了label:bug
TorchNPU-BotTorchNPU-Bot成员
17 天前 添加了label:bot-triaged;删除了label:triage-review
TorchNPU-Bot
TorchNPU-Bot成员
17 天前 评论:

检测到当前 issue 已关联 PR,自动添加标签:bot-triaged

likedislike