已关闭
[Bug]: OperationFuser 在 no-grad 上下文中对非叶子张量调用 requires_grad_ 导致 RuntimeError #22
Fanlu创建于  8月4日关闭于  4 天前
Fanlu成员
8月4日 创建

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

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

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

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

环境信息

例如:
- 操作系统
- 昇腾硬件信息
- CANN软件版本
- 安装的对应软件版本

原始报错信息:

[rank0]:   File "/home/xxx/workspace/Megatron-LM/megatron/core/transformer/transformer_layer.py", line 626, in _forward_attention
[rank0]:     input_layernorm_output = apply_module(self.input_layernorm)(hidden_states)
[rank0]:                              ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank0]:   File "/home/miniconda3/envs/fanlu_py312/lib/python3.12/site-packages/torch/nn/modules/module.py", line 1751, in _wrapped_call_impl
[rank0]:     return self._call_impl(*args, **kwargs)
[rank0]:            ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank0]:   File "/home/miniconda3/envs/fanlu_py312/lib/python3.12/site-packages/torch/nn/modules/module.py", line 1762, in _call_impl
[rank0]:     return forward_call(*args, **kwargs)
[rank0]:            ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank0]:   File "/home/xxx/workspace/Megatron-LM/megatron/core/extensions/transformer_engine.py", line 790, in forward
[rank0]:     return self._fused_impl[0](hidden_states)
[rank0]:            ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank0]:   File "/home/miniconda3/envs/fanlu_py312/lib/python3.12/site-packages/torch/nn/modules/module.py", line 1751, in _wrapped_call_impl
[rank0]:     return self._call_impl(*args, **kwargs)
[rank0]:            ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank0]:   File "/home/miniconda3/envs/fanlu_py312/lib/python3.12/site-packages/torch/nn/modules/module.py", line 1762, in _call_impl
[rank0]:     return forward_call(*args, **kwargs)
[rank0]:            ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank0]:   File "/home/xxx/workspace/TransformerEngineNPU/transformer_engine/pytorch/ops/sequential.py", line 189, in forward
[rank0]:     xs = module_group(*xs)
[rank0]:          ^^^^^^^^^^^^^^^^^
[rank0]:   File "/home/xxx/workspace/TransformerEngineNPU/transformer_engine/pytorch/ops/fuser.py", line 510, in __call__
[rank0]:     return forward_func(*args)
[rank0]:            ^^^^^^^^^^^^^^^^^^^
[rank0]:   File "/home/xxx/workspace/TransformerEngineNPU/transformer_engine/pytorch/ops/fuser.py", line 134, in forward
[rank0]:     y.requires_grad_(idx >= fuser.first_op_requiring_backward)
[rank0]: RuntimeError: you can only change requires_grad flags of leaf variables. If you want to use a computed variable in a subgraph that doesn't require differentiation use var_no_grad = var.detach(). 

🐛 问题描述

OperationFuser 在 no-grad 上下文中对非叶子张量调用 requires_grad_ 导致 RuntimeError

问题背景

OperationFusertransformer_engine/pytorch/ops/fuser.py)负责将多个 BasicOperation
融合成一条流水线执行。Megatron-LM 的 --fused-residual-rmsnorm 特性
TEFusedResidualRMSNorm)基于它构建了 Sequential([MakeExtraOutput, RMSNorm])
流水线,用于将 "残差相加 + RMSNorm" 融合,其中 MakeExtraOutput 将输入张量原样透传,
作为额外的输出(即残差)。

当该融合算子在 无梯度(no-grad)上下文 中被调用,且输入是 非叶子张量 时,
fuser 的前向会在设置额外输出的 requires_grad 标志时抛出 RuntimeError,导致训练直接崩溃。

实际触发场景:Megatron-LM 的 MLA(Multi-Latent Attention)layer spec 中,
input_layernorm 使用独立的 layer_norm(has_residual=True)(走 TEFusedResidualRMSNorm),
配合 --recompute-granularity full 激活重计算时必现。

报错信息

RuntimeError: you can only change requires_grad flags of leaf variables.
If you want to use a computed variable in a subgraph that doesn't require
differentiation use var_no_grad = var.detach().

调用栈(自底向上关键帧):

transformer_layer.py            _forward_attention
  -> input_layernorm(hidden_states)                       # TEFusedResidualRMSNorm
extensions/transformer_engine.py  TEFusedResidualRMSNorm.forward
  -> self._fused_impl[0](hidden_states)                   # Sequential ops
ops/sequential.py               SequentialOps.forward
  -> module_group(*xs)                                    # OperationFuser
ops/fuser.py                    OperationFuser.__call__
  -> forward_func(*args)                                  # no-grad 分支直接调 forward
ops/fuser.py                    _OperationFuserAutogradFunction.forward
  -> y.requires_grad_(idx >= fuser.first_op_requiring_backward)   # 崩溃点

触发条件

同时满足以下条件时必现:

  1. ops 流水线中包含会透传/派生额外输出的算子(如 MakeExtraOutput),
    其额外输出直接或间接来自 fuser 的输入张量;
  2. fuser 在 torch.is_grad_enabled() == False 的上下文中被调用
    (例如 Megatron-LM 的 CheckpointFunction.forwardtorch.no_grad()
    内执行被重计算块包裹的前向,见 megatron/core/tensor_parallel/random.py);
  3. fuser 的输入张量是非叶子张量(训练时的中间计算结果,如 embedding 输出);
  4. PyTorch 语义:对非叶子张量调用 requires_grad_(False) 会抛出 RuntimeError
    (对非叶子张量调用 requires_grad_(True) 是合法的 no-op,不会报错)。

Megatron-LM 侧的复现参数组合:

--multi-latent-attention
--fused-residual-rmsnorm
--recompute-granularity full
--recompute-method block

说明:非 MLA 的 GPT layer spec 将 layernorm 融合进 LayerNormLinear
column_parallel_layer_norm_linear),没有独立的 input_layernorm
不走 TEFusedResidualRMSNorm,因此不受影响——这解释了为什么同样的
recompute 配置下只有 MLA 配置会崩溃。

不依赖 NPU 的最小复现

PyTorch 语义层面的最小复现(fuser 崩溃点的等价逻辑):

import torch

x = torch.randn(4, requires_grad=True)
h = x * 2                      # 非叶子张量,requires_grad=True
with torch.no_grad():
    h.requires_grad_(False)    # RuntimeError: you can only change
                               # requires_grad flags of leaf variables.

根因分析

以本仓库 2.13.0 版本代码为准:

  1. OperationFuser.__call__(fuser.py)检测到 is_grad_enabled=False 时,
    不走 autograd Function 的 apply,而是直接调用
    _OperationFuserAutogradFunction.forward,并把 first_op_requiring_backward
    置为算子总数(即认为所有算子都不需要反向);

  2. 前向逐算子执行,MakeExtraOutput.fuser_forward 的实现是
    return input_, [(input_,)] —— 将输入张量原样作为额外输出返回;

  3. 随后 fuser 无条件执行:

    for idx, ys in zip(basic_op_idxs, fused_op_extra_outputs):
        for y in ys:
            y.requires_grad_(idx >= fuser.first_op_requiring_backward)
        extra_outputs[idx] = ys
    

    此时 idx=0 < first_op_requiring_backward=算子总数,即对 y(也就是
    非叶子的 input_)调用 requires_grad_(False),触发 RuntimeError;

  4. 前向末尾的 x.requires_grad_(fuser.first_op_requiring_backward < ...)
    存在同样的问题(当最终输出恰好是透传的非叶子张量时);

  5. 为什么 grad-enabled 路径不崩:此时 first_op_requiring_backward=0
    目标是 requires_grad_(True),而 PyTorch 对非叶子张量设置 True 是
    合法的 no-op。因此该 bug 只在 no-grad 上下文(如激活重计算的首次前向)暴露。

上游 NVIDIA/TransformerEngine 的 main 分支已经修复了该问题:
_OperationFuserAutogradFunction.forward 中引入 set_output_requires_grad
参数,对两处 requires_grad_ 调用加了保护,上游代码注释明确写道:

Note: We call forward directly when is_grad_enabled=False,
which can expose non-leaf tensors to the inner ops. Avoid
problems in this case by passing set_output_requires_grad=False.

本仓库(2.13.0)尚未包含该修复。

相关代码

本仓库(TransformerEngineNPU 2.13.0):

  • transformer_engine/pytorch/ops/fuser.py
    • OperationFuser.__call__:no-grad 分支直接调用 forward
      并设置 first_op_requiring_backward = self._num_basic_ops
    • _OperationFuserAutogradFunction.forward
      • 额外输出处:y.requires_grad_(idx >= fuser.first_op_requiring_backward)
      • 主输出处:x.requires_grad_(fuser.first_op_requiring_backward < fuser._num_basic_ops)
  • transformer_engine/pytorch/ops/basic/extra_input_output.py
    • MakeExtraOutput.fuser_forwardreturn input_, [(input_,)](透传输入)。

上游参照(NVIDIA/TransformerEngine main):

  • transformer_engine/pytorch/ops/fuser.pyset_output_requires_grad 相关改动。

Megatron-LM 侧(触发路径,非缺陷方):

  • megatron/core/extensions/transformer_engine.pyTEFusedResidualRMSNorm
    _make_fused_impl 构建 Sequential([MakeExtraOutput, RMSNorm])
  • megatron/core/models/gpt/gpt_layer_specs.py:MLA spec 中
    input_layernorm=backend.layer_norm(has_residual=True)
  • megatron/core/tensor_parallel/random.pyCheckpointFunction.forward
    torch.no_grad() 中执行重计算块的前向。

建议修复方案

Backport 上游 NVIDIA/TransformerEngine main 分支的修复:为
_OperationFuserAutogradFunction.forward 增加 set_output_requires_grad
参数,no-grad 路径下跳过所有 requires_grad_ 调用。该修复已在本地应用并验证
(编译通过,_OperationFuserAutogradFunction 无其他调用点),diff 如下:

--- a/transformer_engine/pytorch/ops/fuser.py
+++ b/transformer_engine/pytorch/ops/fuser.py
@@ _OperationFuserAutogradFunction.forward
     def forward(
         func_ctx: Optional[torch.autograd.function.FunctionCtx],
         input_: torch.Tensor,
         fuser: OperationFuser,
         basic_op_kwargs: list[dict[str, Any]],
+        set_output_requires_grad: bool,
         *params_and_extra_inputs: torch.Tensor,
     ) -> torch.Tensor | tuple[torch.Tensor, ...]:
@@ 前向额外输出
             for idx, ys in zip(basic_op_idxs, fused_op_extra_outputs):
                 for y in ys:
-                    y.requires_grad_(idx >= fuser.first_op_requiring_backward)
+                    if set_output_requires_grad:
+                        y.requires_grad_(idx >= fuser.first_op_requiring_backward)
                 extra_outputs[idx] = ys
@@ 前向主输出
-        x.requires_grad_(fuser.first_op_requiring_backward < fuser._num_basic_ops)
+        if set_output_requires_grad:
+            x.requires_grad_(fuser.first_op_requiring_backward < fuser._num_basic_ops)
@@ backward 返回值
         return (
             dx,  # input_
             None,  # fuser
             None,  # basic_op_kwargs
+            None,  # set_output_requires_grad
             *grad_params_flat,
             *grad_extra_inputs_flat,
         )
@@ OperationFuser.__call__
         # Fuser forward pass
+        # Note: We call forward directly when is_grad_enabled=False,
+        # which can expose non-leaf tensors to the inner ops. Avoid
+        # problems in this case by passing set_output_requires_grad=False.
         if is_grad_enabled:
             forward_func = _OperationFuserAutogradFunction.apply
             args = []
         else:
             forward_func = _OperationFuserAutogradFunction.forward
             args = [None]
         args += (
             input,
             self,
             basic_op_kwargs,
+            is_grad_enabled,  # set_output_requires_grad
             *self._flat_basic_op_params,
             *extra_inputs,
         )
         return forward_func(*args)

修复的语义正确性说明:no-grad 路径下不需要也不应该修改输出张量的
requires_grad 标志;透传的非叶子张量本身的 requires_grad 状态由
autograd 图决定,无需(也不能)在这里干预。

环境信息

  • TransformerEngineNPU: 2.13.0
  • PyTorch: 2.13.0
  • Megatron-LM: --fused-residual-rmsnorm + MLA + full recompute 组合

临时规避

在修复合入前,使用方可从训练参数中移除 --fused-residual-rmsnorm
(MLA 的 input_layernorm 会退回普通 te.RMSNorm,功能正确但失去该融合优化)。

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

likedislike
FFanlu成员
8月4日 添加了label:bug
FFanlu成员
6 天前 关联了pull request:fix(backend): backport set_output_requires_grad guard to fuser no-grad path
ascend-robotascend-robot成员
4 天前 关闭了 issue
ascend-robotascend-robot成员
4 天前 issue状态由 TODO 改变为 DONE
ascend-robotascend-robot成员
4 天前 添加了label:resolved