From ed22445edc1311eff673a9d117d7f1493313fd85 Mon Sep 17 00:00:00 2001 From: Chase Block Date: Fri, 7 Aug 2026 13:43:50 -0700 Subject: [PATCH 1/6] 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/6] [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/6] 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/6] 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], ...]: From 9e19715810d657648b2d936c1dbdc39bb8c9f908 Mon Sep 17 00:00:00 2001 From: Chase Block Date: Tue, 11 Aug 2026 07:36:38 -0700 Subject: [PATCH 5/6] Move linear backward function to fused_mla_q_uproj.py Signed-off-by: Chase Block --- .../pytorch/attention/fused_mla_q_uproj.py | 54 +++++++++--- transformer_engine/pytorch/module/linear.py | 85 ------------------- 2 files changed, 40 insertions(+), 99 deletions(-) diff --git a/transformer_engine/pytorch/attention/fused_mla_q_uproj.py b/transformer_engine/pytorch/attention/fused_mla_q_uproj.py index d09d2871e9..d13a7baf96 100644 --- a/transformer_engine/pytorch/attention/fused_mla_q_uproj.py +++ b/transformer_engine/pytorch/attention/fused_mla_q_uproj.py @@ -7,6 +7,7 @@ from __future__ import annotations import functools import os +import weakref from importlib.metadata import PackageNotFoundError, version as get_pkg_version import torch @@ -14,6 +15,7 @@ from packaging.version import Version as PkgVersion from ..constants import MXFP8_BLOCK_SCALING_SIZE +from ..distributed import get_distributed_world_size from ..quantized_tensor import QuantizedTensor from ..tensor.mxfp8_tensor import MXFP8Quantizer, MXFP8Tensor from ..utils import get_device_compute_capability @@ -153,23 +155,47 @@ def backward_linear( 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, + """Linear backward for the fused Q up-proj.""" + from ..module.linear import LinearBwdArgs, _linear_backward, _2X_ACC_DGRAD, _2X_ACC_WGRAD + + 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, + grad_output_quantizer=grad_output_quantizer, + 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="column", + tp_group=tp_group, + tp_size=tp_size, + tensor_parallel=tp_size > 1, + sequence_parallel=sequence_parallel, + 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, grad_bias = _linear_backward(bwd_args) + return dgrad, wgrad, grad_bias + @classmethod def wrap_mxfp8( cls, diff --git a/transformer_engine/pytorch/module/linear.py b/transformer_engine/pytorch/module/linear.py index 1950096fb2..99b321826b 100644 --- a/transformer_engine/pytorch/module/linear.py +++ b/transformer_engine/pytorch/module/linear.py @@ -786,91 +786,6 @@ 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], 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, 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) - - 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, grad_bias = _linear_backward(bwd_args) - return dgrad, wgrad, grad_bias - def _linear_backward(args: LinearBwdArgs) -> Tuple[Union[torch.Tensor, None], ...]: """Backward implementation for the linear layer. From 9da364fec53b498542751211dab17177d4b3139e Mon Sep 17 00:00:00 2001 From: "pre-commit-ci[bot]" <66853113+pre-commit-ci[bot]@users.noreply.github.com> Date: Tue, 11 Aug 2026 14:40:10 +0000 Subject: [PATCH 6/6] [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci --- transformer_engine/pytorch/module/linear.py | 1 - 1 file changed, 1 deletion(-) diff --git a/transformer_engine/pytorch/module/linear.py b/transformer_engine/pytorch/module/linear.py index 99b321826b..56622db5e6 100644 --- a/transformer_engine/pytorch/module/linear.py +++ b/transformer_engine/pytorch/module/linear.py @@ -786,7 +786,6 @@ def _linear_setup_ctx( return (saved_inputmat, wt_save, saved_weight, saved_bias) - def _linear_backward(args: LinearBwdArgs) -> Tuple[Union[torch.Tensor, None], ...]: """Backward implementation for the linear layer.