[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().
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().
import torch
x = torch.randn(4, requires_grad=True)
h = x * 2# 非叶子张量,requires_grad=Truewith torch.no_grad():
h.requires_grad_(False) # RuntimeError: you can only change# requires_grad flags of leaf variables.
for idx, ys inzip(basic_op_idxs, fused_op_extra_outputs):
for y in ys:
y.requires_grad_(idx >= fuser.first_op_requiring_backward)
extra_outputs[idx] = ys
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.
在提交新问题之前,请确保您已经在社区中搜索过相关问题,并使用了社区中提供的资源/工具后,仍未找到满意的解决方式。
⚠️ 安全信息提醒:请仔细检查提供的文本内容,确保其不包含敏感数据信息,包括但不限于:
在分享配置信息或代码示例时,请将敏感信息脱敏处理,或使用
<TOKEN>等占位符替代原有内容。环境信息
原始报错信息:
🐛 问题描述
OperationFuser 在 no-grad 上下文中对非叶子张量调用 requires_grad_ 导致 RuntimeError
问题背景
OperationFuser(transformer_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激活重计算时必现。报错信息
调用栈(自底向上关键帧):
触发条件
同时满足以下条件时必现:
MakeExtraOutput),其额外输出直接或间接来自 fuser 的输入张量;
torch.is_grad_enabled() == False的上下文中被调用(例如 Megatron-LM 的
CheckpointFunction.forward在torch.no_grad()内执行被重计算块包裹的前向,见
megatron/core/tensor_parallel/random.py);requires_grad_(False)会抛出 RuntimeError(对非叶子张量调用
requires_grad_(True)是合法的 no-op,不会报错)。Megatron-LM 侧的复现参数组合:
说明:非 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 版本代码为准:
OperationFuser.__call__(fuser.py)检测到is_grad_enabled=False时,不走 autograd Function 的
apply,而是直接调用_OperationFuserAutogradFunction.forward,并把first_op_requiring_backward置为算子总数(即认为所有算子都不需要反向);
前向逐算子执行,
MakeExtraOutput.fuser_forward的实现是return input_, [(input_,)]—— 将输入张量原样作为额外输出返回;随后 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;前向末尾的
x.requires_grad_(fuser.first_op_requiring_backward < ...)存在同样的问题(当最终输出恰好是透传的非叶子张量时);
为什么 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_调用加了保护,上游代码注释明确写道:本仓库(2.13.0)尚未包含该修复。
相关代码
本仓库(TransformerEngineNPU 2.13.0):
transformer_engine/pytorch/ops/fuser.pyOperationFuser.__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.pyMakeExtraOutput.fuser_forward:return input_, [(input_,)](透传输入)。上游参照(NVIDIA/TransformerEngine main):
transformer_engine/pytorch/ops/fuser.py:set_output_requires_grad相关改动。Megatron-LM 侧(触发路径,非缺陷方):
megatron/core/extensions/transformer_engine.py:TEFusedResidualRMSNorm,_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.py:CheckpointFunction.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 图决定,无需(也不能)在这里干预。
环境信息
--fused-residual-rmsnorm+ MLA + full recompute 组合临时规避
在修复合入前,使用方可从训练参数中移除
--fused-residual-rmsnorm(MLA 的
input_layernorm会退回普通te.RMSNorm,功能正确但失去该融合优化)。欢迎加入社区,感谢您对社区的贡献 🎉!