From ed22445edc1311eff673a9d117d7f1493313fd85 Mon Sep 17 00:00:00 2001 From: Chase Block Date: Fri, 7 Aug 2026 13:43:50 -0700 Subject: [PATCH 1/4] Add a backwards linear function to be used with the fused mla q up-proj Signed-off-by: Chase Block --- .../pytorch/attention/fused_mla_q_uproj.py | 8 ++ transformer_engine/pytorch/module/linear.py | 87 +++++++++++++++++++ 2 files changed, 95 insertions(+) diff --git a/transformer_engine/pytorch/attention/fused_mla_q_uproj.py b/transformer_engine/pytorch/attention/fused_mla_q_uproj.py index c176985254..4d8ab3c207 100644 --- a/transformer_engine/pytorch/attention/fused_mla_q_uproj.py +++ b/transformer_engine/pytorch/attention/fused_mla_q_uproj.py @@ -142,6 +142,14 @@ def run( # 2nd return is the activation to save for wgrad: MXFP8 (fp8 path) or bf16 (16-bit path). return query, x_saved + @classmethod + def backward_linear(cls, grad_output, x_saved, w_q, act_dtype, wgrad_store, + fuse_wgrad_accumulation, tp_group, sequence_parallel, **kwargs): + """Linear backward for the fused Q up-proj — delegates to :func:`~transformer_engine.pytorch.module.linear.backward_linear`.""" + from ..module.linear import backward_linear as _bwd + return _bwd(grad_output, x_saved, w_q, act_dtype, wgrad_store, + fuse_wgrad_accumulation, tp_group, sequence_parallel, **kwargs) + @classmethod def wrap_mxfp8( cls, diff --git a/transformer_engine/pytorch/module/linear.py b/transformer_engine/pytorch/module/linear.py index 56622db5e6..be793e3f23 100644 --- a/transformer_engine/pytorch/module/linear.py +++ b/transformer_engine/pytorch/module/linear.py @@ -786,6 +786,93 @@ def _linear_setup_ctx( return (saved_inputmat, wt_save, saved_weight, saved_bias) +def backward_linear( + grad_output: torch.Tensor, + x_saved, + w_q, + act_dtype: torch.dtype, + wgrad_store, + fuse_wgrad_accumulation: bool, + tp_group, + sequence_parallel: bool, + *, + use_bias: bool = False, + requires_dgrad: bool = True, + requires_wgrad: bool = True, + parallel_mode: str = "column", + backward_input_needs_gather: bool = False, +) -> Tuple[Optional[torch.Tensor], Optional[torch.Tensor]]: + """Linear backward for fused operations that bypass TE's autograd chain. + + Wraps :func:`_linear_backward` with a simplified interface for callers + (e.g. Megatron's fused MLA Q up-proj) that run their own forward kernel + and need to delegate the projection backward to TE. + + Args: + grad_output: upstream gradient (e.g. post-RoPE-backward) ``[tokens, out_features]``. + x_saved: activation saved from the forward (``MXFP8Tensor`` or bf16). + w_q: weight (``MXFP8Tensor`` for FP8 path, bf16 tensor otherwise). + act_dtype: output dtype for the dgrad tensor. + wgrad_store: optional deferred weight-grad store. + fuse_wgrad_accumulation: accumulate wgrad directly into ``w_q.main_grad``. + tp_group: tensor-parallel process group (or ``None``). + sequence_parallel: whether sequence parallelism is active. + use_bias: compute a bias gradient (default ``False``). + requires_dgrad: compute dgrad (default ``True``). + requires_wgrad: compute wgrad (default ``True``). + parallel_mode: cuBLAS parallel mode (default ``"column"``). + backward_input_needs_gather: all-gather ``x_saved`` before the wgrad + GEMM (default ``False`` — assumes fused forward pre-gathers). + + Returns: + ``(dgrad, wgrad)`` — ``wgrad`` is a typed dummy when + ``fuse_wgrad_accumulation=True``. + """ + import weakref + + tp_size = get_distributed_world_size(tp_group) if tp_group is not None else 1 + fp8 = isinstance(w_q, QuantizedTensor) + + grad_output_quantizer = None + if fp8: + grad_output_quantizer = MXFP8Quantizer( + fp8_dtype=tex.DType.kFloat8E4M3, rowwise=True, columnwise=True + ) + grad_output_quantizer.optimize_for_gemm = True + + bwd_args = LinearBwdArgs( + grad_output=grad_output, + inputmat=x_saved, + weight_fp8=w_q, + saved_weight=w_q, + bias=None, + grad_output_quantizer=grad_output_quantizer, + use_bias=use_bias, + requires_dgrad=requires_dgrad, + requires_wgrad=requires_wgrad, + inp_shape=x_saved.shape, + activation_dtype=act_dtype, + fp8=fp8, + dgrad_use_split_accumulator=_2X_ACC_DGRAD, + wgrad_use_split_accumulator=_2X_ACC_WGRAD, + is_weight_param_quantized=fp8, + parallel_mode=parallel_mode, + tp_group=tp_group, + tp_size=tp_size, + tensor_parallel=tp_size > 1, + sequence_parallel=sequence_parallel, + backward_input_needs_gather=backward_input_needs_gather, + is_fsdp2=False, + fuse_wgrad_accumulation=fuse_wgrad_accumulation, + wgrad_store=wgrad_store, + origin_weight_ref=weakref.ref(w_q) if fuse_wgrad_accumulation else None, + main_grad_func=(lambda: w_q.main_grad) if fuse_wgrad_accumulation else None, + ) + + wgrad, dgrad, _ = _linear_backward(bwd_args) + return dgrad, wgrad + + def _linear_backward(args: LinearBwdArgs) -> Tuple[Union[torch.Tensor, None], ...]: """Backward implementation for the linear layer. From 346249c26e7c7c182375a9e66e98ef574e004bbe Mon Sep 17 00:00:00 2001 From: "pre-commit-ci[bot]" <66853113+pre-commit-ci[bot]@users.noreply.github.com> Date: Fri, 7 Aug 2026 21:16:08 +0000 Subject: [PATCH 2/4] [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci --- .../pytorch/attention/fused_mla_q_uproj.py | 28 ++++++++++++++++--- 1 file changed, 24 insertions(+), 4 deletions(-) diff --git a/transformer_engine/pytorch/attention/fused_mla_q_uproj.py b/transformer_engine/pytorch/attention/fused_mla_q_uproj.py index 4d8ab3c207..d09d2871e9 100644 --- a/transformer_engine/pytorch/attention/fused_mla_q_uproj.py +++ b/transformer_engine/pytorch/attention/fused_mla_q_uproj.py @@ -143,12 +143,32 @@ def run( return query, x_saved @classmethod - def backward_linear(cls, grad_output, x_saved, w_q, act_dtype, wgrad_store, - fuse_wgrad_accumulation, tp_group, sequence_parallel, **kwargs): + def backward_linear( + cls, + grad_output, + x_saved, + w_q, + act_dtype, + wgrad_store, + fuse_wgrad_accumulation, + tp_group, + sequence_parallel, + **kwargs, + ): """Linear backward for the fused Q up-proj — delegates to :func:`~transformer_engine.pytorch.module.linear.backward_linear`.""" from ..module.linear import backward_linear as _bwd - return _bwd(grad_output, x_saved, w_q, act_dtype, wgrad_store, - fuse_wgrad_accumulation, tp_group, sequence_parallel, **kwargs) + + return _bwd( + grad_output, + x_saved, + w_q, + act_dtype, + wgrad_store, + fuse_wgrad_accumulation, + tp_group, + sequence_parallel, + **kwargs, + ) @classmethod def wrap_mxfp8( From 554d64438ffa659234126c6097e825948e3ade58 Mon Sep 17 00:00:00 2001 From: Chase Block Date: Fri, 7 Aug 2026 14:26:05 -0700 Subject: [PATCH 3/4] Remove redundant import in backward_linear Signed-off-by: Chase Block --- transformer_engine/pytorch/module/linear.py | 2 -- 1 file changed, 2 deletions(-) diff --git a/transformer_engine/pytorch/module/linear.py b/transformer_engine/pytorch/module/linear.py index be793e3f23..a577ca7c3e 100644 --- a/transformer_engine/pytorch/module/linear.py +++ b/transformer_engine/pytorch/module/linear.py @@ -828,8 +828,6 @@ def backward_linear( ``(dgrad, wgrad)`` — ``wgrad`` is a typed dummy when ``fuse_wgrad_accumulation=True``. """ - import weakref - tp_size = get_distributed_world_size(tp_group) if tp_group is not None else 1 fp8 = isinstance(w_q, QuantizedTensor) From 050a4ed07aeac1209a11b752e340d8a0440a602b Mon Sep 17 00:00:00 2001 From: Chase Block Date: Mon, 10 Aug 2026 08:06:05 -0700 Subject: [PATCH 4/4] Handle biad gradient in lin bwd wrapper. Signed-off-by: Chase Block --- transformer_engine/pytorch/module/linear.py | 11 ++++++----- 1 file changed, 6 insertions(+), 5 deletions(-) diff --git a/transformer_engine/pytorch/module/linear.py b/transformer_engine/pytorch/module/linear.py index a577ca7c3e..1950096fb2 100644 --- a/transformer_engine/pytorch/module/linear.py +++ b/transformer_engine/pytorch/module/linear.py @@ -801,7 +801,7 @@ def backward_linear( requires_wgrad: bool = True, parallel_mode: str = "column", backward_input_needs_gather: bool = False, -) -> Tuple[Optional[torch.Tensor], Optional[torch.Tensor]]: +) -> Tuple[Optional[torch.Tensor], Optional[torch.Tensor], Optional[torch.Tensor]]: """Linear backward for fused operations that bypass TE's autograd chain. Wraps :func:`_linear_backward` with a simplified interface for callers @@ -825,8 +825,9 @@ def backward_linear( GEMM (default ``False`` — assumes fused forward pre-gathers). Returns: - ``(dgrad, wgrad)`` — ``wgrad`` is a typed dummy when - ``fuse_wgrad_accumulation=True``. + ``(dgrad, wgrad, grad_bias)`` — ``wgrad`` is a typed dummy when + ``fuse_wgrad_accumulation=True``; ``grad_bias`` is ``None`` when + ``use_bias=False``. """ tp_size = get_distributed_world_size(tp_group) if tp_group is not None else 1 fp8 = isinstance(w_q, QuantizedTensor) @@ -867,8 +868,8 @@ def backward_linear( main_grad_func=(lambda: w_q.main_grad) if fuse_wgrad_accumulation else None, ) - wgrad, dgrad, _ = _linear_backward(bwd_args) - return dgrad, wgrad + wgrad, dgrad, grad_bias = _linear_backward(bwd_args) + return dgrad, wgrad, grad_bias def _linear_backward(args: LinearBwdArgs) -> Tuple[Union[torch.Tensor, None], ...]: