From a4d79a71f870f8ab9bef08cf27701b4de68a8157 Mon Sep 17 00:00:00 2001 From: Ravi Ghadia Date: Thu, 27 Aug 2026 17:19:00 -0700 Subject: [PATCH] Remove redundant runtime checks for activation recompute in MLP from _ScaledUnary class in activation.py Signed-off-by: Ravi Ghadia --- transformer_engine/pytorch/ops/basic/activation.py | 12 ------------ 1 file changed, 12 deletions(-) diff --git a/transformer_engine/pytorch/ops/basic/activation.py b/transformer_engine/pytorch/ops/basic/activation.py index a974d41ef9..7530825c40 100644 --- a/transformer_engine/pytorch/ops/basic/activation.py +++ b/transformer_engine/pytorch/ops/basic/activation.py @@ -402,12 +402,6 @@ def fuser_forward( next_op_input_quantizer: Optional[Quantizer], # pylint: disable=unused-argument basic_op_kwargs: list[dict[str, Any]], # pylint: disable=unused-argument ) -> tuple[torch.Tensor, Sequence[Sequence[torch.Tensor]]]: - if self.activation_recompute_in_mlp: - raise RuntimeError( - f"{self.__class__.__name__}(activation_recompute_in_mlp=True) requires the " - "fused grouped MLP path." - ) - extra_input = basic_op_extra_inputs[0][0] if torch.is_autocast_enabled(): @@ -445,12 +439,6 @@ def fuser_backward( ]: del basic_op_grad_extra_outputs - if self.activation_recompute_in_mlp: - raise RuntimeError( - f"{self.__class__.__name__}(activation_recompute_in_mlp=True) requires the " - "fused grouped MLP path." - ) - ctx = basic_op_ctxs[0] x, scales = ctx.saved_tensors x = maybe_dequantize(x.contiguous(), ctx.dtype)