diff --git a/megatron/core/extensions/transformer_engine.py b/megatron/core/extensions/transformer_engine.py index 9393ad825c3..337555707cb 100644 --- a/megatron/core/extensions/transformer_engine.py +++ b/megatron/core/extensions/transformer_engine.py @@ -614,7 +614,16 @@ def __new__(cls, config: TransformerConfig): activation_func_kwargs = {} if config.activation_func_fp8_input_store: activation_func_kwargs["cache_quantized_input"] = True - layer = layer_type(**activation_func_kwargs) + if config.use_situ_glu: + from megatron.core.fusions.cutedsl_situ_glu import make_situ_glu + + layer = make_situ_glu( + beta1=config.situ_glu_beta1, + beta2=config.situ_glu_beta2, + **activation_func_kwargs, + ) + else: + layer = layer_type(**activation_func_kwargs) return layer else: diff --git a/megatron/core/fusions/cutedsl_situ_glu.py b/megatron/core/fusions/cutedsl_situ_glu.py new file mode 100644 index 00000000000..612783ca717 --- /dev/null +++ b/megatron/core/fusions/cutedsl_situ_glu.py @@ -0,0 +1,543 @@ +# Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +"""SiTU-GLU activation builders with an MCore-local CuTe DSL fallback.""" + +from __future__ import annotations + +import math +from collections.abc import Sequence +from functools import lru_cache +from typing import Any, Optional + +import torch +from transformer_engine.pytorch.cpu_offload import is_cpu_offload_enabled, mark_activation_offload +from transformer_engine.pytorch.ops._common import maybe_dequantize +from transformer_engine.pytorch.ops.op import BasicOperation, OperationContext +from transformer_engine.pytorch.utils import clear_tensor_data + + +def situ_glu_reference( + input_: torch.Tensor, beta1: float = 4.0, beta2: float = 25.0 +) -> torch.Tensor: + """Compute SiTU-GLU with the gate in the first half of the last dimension.""" + gate, up = input_.chunk(2, dim=-1) + return (beta1 * torch.tanh(gate / beta1) * torch.sigmoid(gate)) * ( + beta2 * torch.tanh(up / beta2) + ) + + +def _validate_betas(beta1: float, beta2: float) -> tuple[float, float]: + beta1 = float(beta1) + beta2 = float(beta2) + if not math.isfinite(beta1) or beta1 <= 0.0: + raise ValueError(f"SiTU-GLU beta1 must be finite and positive, got {beta1}.") + if not math.isfinite(beta2) or beta2 <= 0.0: + raise ValueError(f"SiTU-GLU beta2 must be finite and positive, got {beta2}.") + return beta1, beta2 + + +try: + import cuda.bindings.driver as cuda + import cutlass + import cutlass.cute as cute + from cutlass.cute.runtime import from_dlpack, make_fake_stream + + _CUTE_AVAILABLE = True +except ImportError: + cuda = None + cutlass = None + cute = None + from_dlpack = None + make_fake_stream = None + _CUTE_AVAILABLE = False + + +if _CUTE_AVAILABLE: + + @cute.kernel + def _situ_glu_forward_kernel( + input_: cute.Tensor, + output: cute.Tensor, + rows: cutlass.Int32, + width: cutlass.Int32, + beta1: cutlass.Float32, + beta2: cutlass.Float32, + inv_beta1: cutlass.Float32, + inv_beta2: cutlass.Float32, + beta_product: cutlass.Float32, + interleave_size: cutlass.Int32, + ): + tidx, _, _ = cute.arch.thread_idx() + bidx, _, _ = cute.arch.block_idx() + index = bidx * 256 + tidx + elements = rows * width + if index < elements: + row = index // width + col = index - row * width + gate_col = col + up_col = col + width + if interleave_size > 0: + block = col // interleave_size + offset = col - block * interleave_size + gate_col = block * 2 * interleave_size + offset + up_col = gate_col + interleave_size + gate = cutlass.Float32(input_[row, gate_col]) + up = cutlass.Float32(input_[row, up_col]) + gate_tanh = cute.math.tanh(gate * inv_beta1, fastmath=True) + up_tanh = cute.math.tanh(up * inv_beta2, fastmath=True) + sigmoid = cutlass.Float32(0.0) + if beta1 == cutlass.Float32(4.0): + reciprocal = cute.arch.rcp_approx(cutlass.Float32(1.0) + gate_tanh * gate_tanh) + sigmoid = cutlass.Float32(0.5) + gate_tanh * reciprocal + else: + sigmoid = cute.arch.rcp_approx( + cutlass.Float32(1.0) + cute.math.exp(-gate, fastmath=True) + ) + output[row, col] = (beta_product * gate_tanh * sigmoid * up_tanh).to( + output.element_type + ) + + @cute.jit + def _situ_glu_forward_launch( + input_: cute.Tensor, + output: cute.Tensor, + rows: cutlass.Int32, + width: cutlass.Int32, + beta1: cutlass.Float32, + beta2: cutlass.Float32, + inv_beta1: cutlass.Float32, + inv_beta2: cutlass.Float32, + beta_product: cutlass.Float32, + interleave_size: cutlass.Int32, + stream: cuda.CUstream, + ): + _situ_glu_forward_kernel( + input_, + output, + rows, + width, + beta1, + beta2, + inv_beta1, + inv_beta2, + beta_product, + interleave_size, + ).launch(grid=(cute.ceil_div(rows * width, 256), 1, 1), block=(256, 1, 1), stream=stream) + + @cute.kernel + def _situ_glu_backward_kernel( + grad_output: cute.Tensor, + input_: cute.Tensor, + grad_input: cute.Tensor, + rows: cutlass.Int32, + width: cutlass.Int32, + beta1: cutlass.Float32, + beta2: cutlass.Float32, + inv_beta1: cutlass.Float32, + inv_beta2: cutlass.Float32, + interleave_size: cutlass.Int32, + ): + tidx, _, _ = cute.arch.thread_idx() + bidx, _, _ = cute.arch.block_idx() + index = bidx * 256 + tidx + elements = rows * width + if index < elements: + row = index // width + col = index - row * width + grad = cutlass.Float32(grad_output[row, col]) + gate_col = col + up_col = col + width + if interleave_size > 0: + block = col // interleave_size + offset = col - block * interleave_size + gate_col = block * 2 * interleave_size + offset + up_col = gate_col + interleave_size + gate = cutlass.Float32(input_[row, gate_col]) + up = cutlass.Float32(input_[row, up_col]) + gate_tanh = cute.math.tanh(gate * inv_beta1, fastmath=True) + up_tanh = cute.math.tanh(up * inv_beta2, fastmath=True) + sigmoid = cutlass.Float32(0.0) + gate_grad = cutlass.Float32(0.0) + if beta1 == cutlass.Float32(4.0): + reciprocal = cute.arch.rcp_approx(cutlass.Float32(1.0) + gate_tanh * gate_tanh) + sigmoid = cutlass.Float32(0.5) + gate_tanh * reciprocal + gate_grad = (cutlass.Float32(1.0) - gate_tanh * gate_tanh) * ( + cutlass.Float32(0.5) + + cutlass.Float32(2.0) * gate_tanh * reciprocal * reciprocal + ) + else: + sigmoid = cute.arch.rcp_approx( + cutlass.Float32(1.0) + cute.math.exp(-gate, fastmath=True) + ) + gate_grad = (cutlass.Float32(1.0) - gate_tanh * gate_tanh) * sigmoid + gate_grad += beta1 * gate_tanh * sigmoid * (cutlass.Float32(1.0) - sigmoid) + gate_value = beta1 * gate_tanh * sigmoid + up_value = beta2 * up_tanh + up_grad = cutlass.Float32(1.0) - up_tanh * up_tanh + grad_input[row, gate_col] = (grad * up_value * gate_grad).to(grad_input.element_type) + grad_input[row, up_col] = (grad * gate_value * up_grad).to(grad_input.element_type) + + @cute.jit + def _situ_glu_backward_launch( + grad_output: cute.Tensor, + input_: cute.Tensor, + grad_input: cute.Tensor, + rows: cutlass.Int32, + width: cutlass.Int32, + beta1: cutlass.Float32, + beta2: cutlass.Float32, + inv_beta1: cutlass.Float32, + inv_beta2: cutlass.Float32, + interleave_size: cutlass.Int32, + stream: cuda.CUstream, + ): + _situ_glu_backward_kernel( + grad_output, + input_, + grad_input, + rows, + width, + beta1, + beta2, + inv_beta1, + inv_beta2, + interleave_size, + ).launch(grid=(cute.ceil_div(rows * width, 256), 1, 1), block=(256, 1, 1), stream=stream) + + +def _gpu_arch() -> str: + capability = torch.cuda.get_device_capability() + arch = {(9, 0): "sm_90a", (10, 0): "sm_100a", (10, 3): "sm_103a"}.get(capability) + if arch is None: + raise RuntimeError( + f"SiTU-GLU CuTe DSL kernels do not support compute capability {capability}." + ) + return arch + + +def _cute_tensor(tensor: torch.Tensor): + tensor = from_dlpack(tensor.detach(), assumed_align=16, enable_tvm_ffi=True) + return tensor.mark_layout_dynamic(leading_dim=1) + + +@lru_cache(maxsize=None) +def _compile_forward(dtype: torch.dtype, device_index: int, rows: int, double_width: int): + with torch.cuda.device(device_index): + fake_input = torch.empty((rows, double_width), device="cuda", dtype=dtype) + fake_output = torch.empty((rows, double_width // 2), device="cuda", dtype=dtype) + return cute.compile( + _situ_glu_forward_launch, + _cute_tensor(fake_input), + _cute_tensor(fake_output), + cutlass.Int32(1), + cutlass.Int32(1), + cutlass.Float32(4.0), + cutlass.Float32(25.0), + cutlass.Float32(0.25), + cutlass.Float32(0.04), + cutlass.Float32(100.0), + cutlass.Int32(0), + make_fake_stream(use_tvm_ffi_env_stream=False), + options=f"--enable-tvm-ffi --gpu-arch {_gpu_arch()}", + ) + + +@lru_cache(maxsize=None) +def _compile_backward(dtype: torch.dtype, device_index: int, rows: int, double_width: int): + with torch.cuda.device(device_index): + fake_grad = torch.empty((rows, double_width // 2), device="cuda", dtype=dtype) + fake_input = torch.empty((rows, double_width), device="cuda", dtype=dtype) + return cute.compile( + _situ_glu_backward_launch, + _cute_tensor(fake_grad), + _cute_tensor(fake_input), + _cute_tensor(fake_input), + cutlass.Int32(1), + cutlass.Int32(1), + cutlass.Float32(4.0), + cutlass.Float32(25.0), + cutlass.Float32(0.25), + cutlass.Float32(0.04), + cutlass.Int32(0), + make_fake_stream(use_tvm_ffi_env_stream=False), + options=f"--enable-tvm-ffi --gpu-arch {_gpu_arch()}", + ) + + +def _validate_input(input_: torch.Tensor, interleave_size: int = 0) -> torch.Tensor: + if not _CUTE_AVAILABLE: + raise RuntimeError("SiTU-GLU requires the nvidia-cutlass-dsl package.") + if not input_.is_cuda: + raise ValueError("SiTU-GLU CuTe DSL kernels require a CUDA tensor.") + if input_.dtype not in (torch.bfloat16, torch.float16): + raise ValueError("SiTU-GLU CuTe DSL kernels support BF16 and FP16 tensors.") + if input_.shape[-1] % 2: + raise ValueError("SiTU-GLU input width must be even.") + if interleave_size < 0: + raise ValueError("SiTU-GLU interleave size must be non-negative.") + if interleave_size and input_.shape[-1] % (2 * interleave_size): + raise ValueError("SiTU-GLU input width must be divisible by twice the interleave size.") + return input_.contiguous().view(-1, input_.shape[-1]) + + +def _situ_glu_forward( + input_: torch.Tensor, beta1: float, beta2: float, interleave_size: int = 0 +) -> torch.Tensor: + input_2d = _validate_input(input_, interleave_size) + rows, double_width = input_2d.shape + output = torch.empty((rows, double_width // 2), device=input_.device, dtype=input_.dtype) + launcher = _compile_forward(input_.dtype, input_.device.index, rows, double_width) + launcher( + input_2d, + output, + cutlass.Int32(rows), + cutlass.Int32(double_width // 2), + cutlass.Float32(beta1), + cutlass.Float32(beta2), + cutlass.Float32(1.0 / beta1), + cutlass.Float32(1.0 / beta2), + cutlass.Float32(beta1 * beta2), + cutlass.Int32(interleave_size), + cuda.CUstream(torch.cuda.current_stream(input_.device).cuda_stream), + ) + return output.view(*input_.shape[:-1], double_width // 2) + + +def _situ_glu_backward( + grad_output: torch.Tensor, + input_: torch.Tensor, + beta1: float, + beta2: float, + interleave_size: int = 0, +) -> torch.Tensor: + input_2d = _validate_input(input_, interleave_size) + grad_output_2d = grad_output.contiguous().view(-1, grad_output.shape[-1]) + grad_input = torch.empty_like(input_2d) + rows, double_width = input_2d.shape + launcher = _compile_backward(input_.dtype, input_.device.index, rows, double_width) + launcher( + grad_output_2d, + input_2d, + grad_input, + cutlass.Int32(rows), + cutlass.Int32(double_width // 2), + cutlass.Float32(beta1), + cutlass.Float32(beta2), + cutlass.Float32(1.0 / beta1), + cutlass.Float32(1.0 / beta2), + cutlass.Int32(interleave_size), + cuda.CUstream(torch.cuda.current_stream(input_.device).cuda_stream), + ) + return grad_input.view_as(input_) + + +class _SiTUGLUFunction(torch.autograd.Function): + @staticmethod + def forward( + ctx, input_: torch.Tensor, beta1: float, beta2: float, interleave_size: int + ) -> torch.Tensor: + """Run SiTU-GLU forward and save its input for the custom backward.""" + ctx.save_for_backward(input_) + ctx.beta1 = beta1 + ctx.beta2 = beta2 + ctx.interleave_size = interleave_size + return _situ_glu_forward(input_, beta1, beta2, interleave_size) + + @staticmethod + def backward(ctx, grad_output: torch.Tensor): + """Compute the input gradient for the custom SiTU-GLU operation.""" + (input_,) = ctx.saved_tensors + return ( + _situ_glu_backward(grad_output, input_, ctx.beta1, ctx.beta2, ctx.interleave_size), + None, + None, + None, + ) + + +class CuTeDSLSiTUGLU(torch.nn.Module): + """Standalone SiTU-GLU used by shared experts and unfused routed experts.""" + + def __init__(self, beta1: float = 4.0, beta2: float = 25.0, interleave_size: int = 0) -> None: + super().__init__() + self.beta1, self.beta2 = _validate_betas(beta1, beta2) + self.interleave_size = int(interleave_size) + if self.interleave_size < 0: + raise ValueError("SiTU-GLU interleave size must be non-negative.") + + def forward(self, input_: torch.Tensor) -> torch.Tensor: + """Apply standalone SiTU-GLU to the input tensor.""" + return _SiTUGLUFunction.apply(input_, self.beta1, self.beta2, self.interleave_size) + + +def _get_te_situ_glu_ops() -> Optional[tuple[type[torch.nn.Module], type[torch.nn.Module]]]: + """Import the complete public Transformer Engine SiTU-GLU interface.""" + try: + from transformer_engine.pytorch.ops import ScaledSiTUGLU, SiTUGLU + except ImportError: + return None + return SiTUGLU, ScaledSiTUGLU + + +class CuTeDSLScaledSiTUGLU(BasicOperation): + """MCore-local scaled SiTU-GLU fallback for TE's operation-fuser interface.""" + + num_extra_inputs: int = 1 + + def __init__( + self, + glu_interleave_size: Optional[int] = None, + *, + activation_recompute_in_mlp: bool = False, + beta1: float = 4.0, + beta2: float = 25.0, + ) -> None: + super().__init__() + if activation_recompute_in_mlp: + raise ValueError( + f"{self.__class__.__name__} does not support activation recomputation " + "in the fused grouped MLP" + ) + self.beta1, self.beta2 = _validate_betas(beta1, beta2) + self.glu_interleave_size = glu_interleave_size + if self.glu_interleave_size is not None: + self.glu_interleave_size = int(self.glu_interleave_size) + if self.glu_interleave_size <= 0: + raise ValueError("SiTU-GLU interleave size must be positive when set.") + + def op_forward(self, *args, **kwargs) -> None: + """Reject direct BasicOperation execution with a hidden scale input.""" + raise RuntimeError( + f"{self.__class__.__name__} expects its row scales as an operation-fuser input." + ) + + def op_backward(self, *args, **kwargs) -> None: + """Reject direct BasicOperation backward with a hidden scale input.""" + raise RuntimeError( + f"{self.__class__.__name__} expects its row scales as an operation-fuser input." + ) + + def fuser_forward( + self, + basic_op_ctxs: list[OperationContext], + input_: torch.Tensor, + *, + basic_op_extra_inputs: Sequence[Sequence[Optional[torch.Tensor]]], + prev_op_grad_output_quantizer: Optional[Any], + next_op_input_quantizer: Optional[Any], + basic_op_kwargs: list[dict[str, Any]], + ) -> tuple[torch.Tensor, Sequence[Sequence[Optional[torch.Tensor]]]]: + """Apply local SiTU-GLU and the routed-token row scales.""" + del prev_op_grad_output_quantizer, next_op_input_quantizer, basic_op_kwargs + scales = basic_op_extra_inputs[0][0] + if torch.is_autocast_enabled(): + dtype = torch.get_autocast_dtype("cuda") + elif isinstance(input_, torch.Tensor): + dtype = input_.dtype + else: + dtype = scales.dtype + input_ = maybe_dequantize(input_, dtype) + scales = maybe_dequantize(scales, dtype) + output = _situ_glu_forward( + input_, self.beta1, self.beta2, int(self.glu_interleave_size or 0) + ) + output = output * scales.unsqueeze(-1) + + ctx = basic_op_ctxs[0] + if ctx.requires_grad: + if is_cpu_offload_enabled(): + mark_activation_offload(input_) + ctx.input_requires_grad = True + ctx.extra_input_requires_grad = scales.requires_grad + ctx.dtype = dtype + ctx.save_for_backward(input_, scales) + return output, [()] + + def fuser_backward( + self, + basic_op_ctxs: list[OperationContext], + grad_output: torch.Tensor, + *, + basic_op_grad_extra_outputs: Sequence[Sequence[Optional[torch.Tensor]]], + ) -> tuple[ + torch.Tensor, + Sequence[Sequence[Optional[torch.Tensor]]], + Sequence[Sequence[Optional[torch.Tensor]]], + ]: + """Differentiate the local activation and optional row scales.""" + del basic_op_grad_extra_outputs + ctx = basic_op_ctxs[0] + input_, scales = ctx.saved_tensors + input_ = maybe_dequantize(input_, ctx.dtype) + scales = maybe_dequantize(scales, ctx.dtype) + grad_output = maybe_dequantize(grad_output, ctx.dtype) + interleave_size = int(self.glu_interleave_size or 0) + grad_input = _situ_glu_backward( + grad_output * scales.unsqueeze(-1), input_, self.beta1, self.beta2, interleave_size + ) + grad_scales = None + if ctx.extra_input_requires_grad: + output = _situ_glu_forward(input_, self.beta1, self.beta2, interleave_size) + grad_scales = torch.linalg.vecdot(output, grad_output) + clear_tensor_data(ctx.saved_tensors[0]) + return grad_input, [()], [(grad_scales,)] + + +def make_situ_glu( + beta1: float = 4.0, + beta2: float = 25.0, + *, + cache_quantized_input: bool = False, + glu_interleave_size: Optional[int] = None, +) -> torch.nn.Module: + """Construct native TE SiTU-GLU or the standalone MCore CuTe fallback.""" + beta1, beta2 = _validate_betas(beta1, beta2) + te_ops = _get_te_situ_glu_ops() + if te_ops is not None: + situ_glu, _ = te_ops + return situ_glu( + beta1=beta1, + beta2=beta2, + cache_quantized_input=cache_quantized_input, + glu_interleave_size=glu_interleave_size, + ) + if cache_quantized_input: + raise RuntimeError( + "activation_func_fp8_input_store for SiTU-GLU requires " + "https://github.com/NVIDIA/TransformerEngine/pull/3402." + ) + return CuTeDSLSiTUGLU(beta1=beta1, beta2=beta2, interleave_size=int(glu_interleave_size or 0)) + + +def make_scaled_situ_glu( + beta1: float = 4.0, + beta2: float = 25.0, + *, + glu_interleave_size: Optional[int] = None, + activation_recompute_in_mlp: bool = False, +) -> torch.nn.Module: + """Construct native TE scaled SiTU-GLU or the MCore CuTe fuser fallback.""" + beta1, beta2 = _validate_betas(beta1, beta2) + te_ops = _get_te_situ_glu_ops() + if te_ops is not None: + _, scaled_situ_glu = te_ops + return scaled_situ_glu( + glu_interleave_size=glu_interleave_size, + activation_recompute_in_mlp=activation_recompute_in_mlp, + beta1=beta1, + beta2=beta2, + ) + return CuTeDSLScaledSiTUGLU( + glu_interleave_size=glu_interleave_size, + activation_recompute_in_mlp=activation_recompute_in_mlp, + beta1=beta1, + beta2=beta2, + ) + + +__all__ = [ + "CuTeDSLSiTUGLU", + "CuTeDSLScaledSiTUGLU", + "make_scaled_situ_glu", + "make_situ_glu", + "situ_glu_reference", +] diff --git a/megatron/core/transformer/moe/experts.py b/megatron/core/transformer/moe/experts.py index 721dc475c6d..7f2453bf854 100644 --- a/megatron/core/transformer/moe/experts.py +++ b/megatron/core/transformer/moe/experts.py @@ -425,6 +425,11 @@ def _is_fused_impl_supported(self) -> bool: ) if not (use_glu_fusion or use_srelu_fusion): return False + if getattr(self.config, "use_situ_glu", False): + if self.config.activation_func != F.silu: + return False + if self.config.moe_mlp_glu_interleave_size != 32: + return False if self.config.activation_func == F.silu: if self.config.activation_func_clamp_value is not None: if not is_te_min_version("2.17.0.dev0"): @@ -537,7 +542,16 @@ def register_grouped_linear_params( # Activation and post-multiply probs (SwiGLU, clamped GLU, or SReLU). glu_interleave = self.config.moe_mlp_glu_interleave_size activation_recompute_in_mlp = bool(getattr(self, "activation_recompute", False)) - if self.config.activation_func == F.silu and self.config.gated_linear_unit: + if getattr(self.config, "use_situ_glu", False): + from megatron.core.fusions.cutedsl_situ_glu import make_scaled_situ_glu + + op = make_scaled_situ_glu( + beta1=self.config.situ_glu_beta1, + beta2=self.config.situ_glu_beta2, + glu_interleave_size=glu_interleave, + activation_recompute_in_mlp=activation_recompute_in_mlp, + ) + elif self.config.activation_func == F.silu and self.config.gated_linear_unit: clamp_value = self.config.activation_func_clamp_value if clamp_value is not None: clamped_glu_kwargs = { diff --git a/megatron/core/transformer/moe/shared_experts.py b/megatron/core/transformer/moe/shared_experts.py index 6eef5dee9ff..befadc8c6d2 100644 --- a/megatron/core/transformer/moe/shared_experts.py +++ b/megatron/core/transformer/moe/shared_experts.py @@ -437,6 +437,11 @@ def _validate_fused_grouped_swiglu(self) -> None: "moe_shared_expert_glu_interleave_size to be set when " "use_grouped_gemm_for_shared_expert=True." ) + if ( + getattr(self.config, "use_situ_glu", False) + and self.config.moe_shared_expert_glu_interleave_size != 32 + ): + raise ValueError("Fused shared-expert SiTU-GLU requires 32-wide GLU interleaving.") if not isinstance(self.linear_fc1, te.pytorch.Linear): raise ValueError( f"{self.__class__.__name__} expects FC1 to be Transformer Engine Linear, " @@ -491,7 +496,15 @@ def _make_fused_grouped_swiglu_ops(self) -> torch.nn.Module: ops.append(op) clamp_value = self.config.activation_func_clamp_value - if clamp_value is None: + if getattr(self.config, "use_situ_glu", False): + from megatron.core.fusions.cutedsl_situ_glu import make_scaled_situ_glu + + activation_op = make_scaled_situ_glu( + beta1=self.config.situ_glu_beta1, + beta2=self.config.situ_glu_beta2, + glu_interleave_size=glu_interleave_size, + ) + elif clamp_value is None: activation_op = te.pytorch.ops.ScaledSwiGLU(glu_interleave_size=glu_interleave_size) else: activation_op = te.pytorch.ops.ScaledClampedQGeGLU( diff --git a/megatron/core/transformer/transformer_config.py b/megatron/core/transformer/transformer_config.py index e783a017056..8ed2eb5c281 100644 --- a/megatron/core/transformer/transformer_config.py +++ b/megatron/core/transformer/transformer_config.py @@ -1170,6 +1170,15 @@ class TransformerConfig(ModelParallelConfig): use_te_activation_func: bool = False """Whether to use ffn activation functions implemented by TransformerEngine""" + use_situ_glu: bool = False + """Use SiTU-GLU in every gated dense, routed-expert, and shared-expert FFN.""" + + situ_glu_beta1: float = 4.0 + """SiTU-GLU gate tanh soft-cap.""" + + situ_glu_beta2: float = 25.0 + """SiTU-GLU up-branch tanh soft-cap.""" + use_te_rng_tracker: bool = False """ Whether to use the TE or MCore version of the RNG tracker. """ @@ -2412,6 +2421,46 @@ def __post_init__(self): "use_te_activation_func to False" ) + if self.use_situ_glu: + if not self.gated_linear_unit: + raise ValueError("use_situ_glu requires gated_linear_unit=True.") + if not self.use_te_activation_func: + raise ValueError("use_situ_glu requires use_te_activation_func=True.") + if self.activation_func != F.silu: + raise ValueError("use_situ_glu requires activation_func=F.silu.") + if self.activation_func_clamp_value is not None: + raise ValueError( + "use_situ_glu is incompatible with activation_func_clamp_value; " + "use situ_glu_beta1 and situ_glu_beta2 for its tanh soft caps." + ) + if self.activation_func_fp8_input_store: + raise ValueError("use_situ_glu does not support activation_func_fp8_input_store.") + if self.fp8 is not None and self.fp8_recipe not in (Fp8Recipe.mxfp8, "mxfp8"): + raise ValueError( + "The CuTe DSL SiTU-GLU implementation supports FP8 only with " + "fp8_recipe='mxfp8'." + ) + if self.fp4 is not None and self.fp4_recipe not in (Fp4Recipe.nvfp4, "nvfp4"): + raise ValueError( + "The CuTe DSL SiTU-GLU implementation supports FP4 only with " + "fp4_recipe='nvfp4'." + ) + if ( + self.num_moe_experts is not None + and self.moe_grouped_gemm + and self.use_transformer_engine_op_fuser + and (self.fp8 is not None or self.fp4 is not None) + and self.moe_mlp_glu_interleave_size != 32 + ): + raise ValueError( + "Fused block-scaled MoE SiTU-GLU requires " + "moe_mlp_glu_interleave_size=32, matching TE's grouped-MLP layout." + ) + if not math.isfinite(self.situ_glu_beta1) or self.situ_glu_beta1 <= 0: + raise ValueError("situ_glu_beta1 must be finite and positive.") + if not math.isfinite(self.situ_glu_beta2) or self.situ_glu_beta2 <= 0: + raise ValueError("situ_glu_beta2 must be finite and positive.") + if self.activation_func_fp8_input_store: if self.activation_func != F.silu or not self.gated_linear_unit: raise ValueError("Storing activation input in FP8 is supported only for SwiGLU.") diff --git a/megatron/training/argument_utils.py b/megatron/training/argument_utils.py index 124083ead83..02130ac843f 100644 --- a/megatron/training/argument_utils.py +++ b/megatron/training/argument_utils.py @@ -311,17 +311,25 @@ def core_transformer_config_from_args(args, config_class=None): kw_args['num_layers_in_last_pipeline_stage']= args.decoder_last_pipeline_num_layers kw_args['fp8_param'] = args.fp8_param_gather kw_args['fp4_param'] = args.fp4_param_gather - if args.swiglu: + use_situ_glu = getattr(args, 'use_situ_glu', False) + if args.swiglu and use_situ_glu: + raise ValueError("--swiglu and --situ-glu select different GLU activations.") + if use_situ_glu: + kw_args['activation_func'] = F.silu + kw_args['gated_linear_unit'] = True + kw_args['use_te_activation_func'] = True + kw_args['bias_activation_fusion'] = False + elif args.swiglu: kw_args['activation_func'] = F.silu kw_args['gated_linear_unit'] = True kw_args['bias_activation_fusion'] = args.bias_swiglu_fusion else: kw_args['bias_activation_fusion'] = args.bias_gelu_fusion if args.squared_relu: - assert not args.swiglu + assert not args.swiglu and not use_situ_glu kw_args['activation_func'] = squared_relu elif args.quick_geglu: - assert not args.swiglu + assert not args.swiglu and not use_situ_glu kw_args['gated_linear_unit'] = True kw_args['activation_func'] = quick_gelu if args.init_method_xavier_uniform: diff --git a/megatron/training/arguments.py b/megatron/training/arguments.py index 3c9984a6b96..a6e8dfeb9a6 100644 --- a/megatron/training/arguments.py +++ b/megatron/training/arguments.py @@ -6,15 +6,15 @@ import dataclasses import json import os -from pathlib import Path import re import types +from pathlib import Path import torch +from megatron.core.msc_utils import MultiStorageClientFeature from megatron.core.rerun_state_machine import RerunStateMachine from megatron.core.transformer import TransformerConfig -from megatron.core.transformer.pipeline_parallel_layer_layout import PipelineParallelLayerLayout from megatron.core.transformer.cuda_graph_config import ( ALLOWED_INFERENCE_SCOPES, get_deprecated_cuda_graph_modules_migration, @@ -23,23 +23,24 @@ validate_deprecated_cuda_graph_modules_migration_inputs, ) from megatron.core.transformer.enums import AttnBackend, CudaGraphModule, InferenceCudaGraphScope +from megatron.core.transformer.pipeline_parallel_layer_layout import PipelineParallelLayerLayout from megatron.core.utils import ( get_torch_version, is_flashinfer_min_version, is_te_min_version, is_torch_min_version, ) +from megatron.training.argument_utils import ( # noqa: F401 # pylint: disable=unused-import + ArgumentGroupFactory, + core_transformer_config_from_args, +) from megatron.training.global_vars import set_global_variables from megatron.training.utils import ( get_device_arch_version, - update_use_dist_ckpt, print_rank_0, + update_use_dist_ckpt, warn_rank_0, ) -from megatron.core.msc_utils import MultiStorageClientFeature - -from megatron.training.argument_utils import ArgumentGroupFactory, core_transformer_config_from_args # noqa: F401 # pylint: disable=unused-import - def add_megatron_arguments(parser: argparse.ArgumentParser): @@ -398,8 +399,9 @@ def validate_args(args, defaults={}): 'Currently only global and local checkpoints are supported' if args.non_persistent_ckpt_type == 'local': try: - from nvidia_resiliency_ext.checkpointing.local.ckpt_managers.local_manager import \ - LocalCheckpointManager + from nvidia_resiliency_ext.checkpointing.local.ckpt_managers.local_manager import ( + LocalCheckpointManager, + ) except ModuleNotFoundError as e: raise RuntimeError('nvidia_resiliency_ext is required for local checkpointing') from e @@ -756,8 +758,10 @@ def validate_args(args, defaults={}): ) from megatron.core.models.hybrid.hybrid_layer_allocation import ( - Symbols, parse_hybrid_pattern, get_hybrid_total_layer_count, + Symbols, + get_hybrid_total_layer_count, get_hybrid_total_pipeline_segment_count, + parse_hybrid_pattern, ) sep = Symbols.MTP_SEPARATOR @@ -1282,8 +1286,17 @@ def validate_args(args, defaults={}): _check_arg_is_not_none(args, req_arg) # Checks. + if args.use_situ_glu: + if args.swiglu or args.quick_geglu or args.squared_relu: + raise ValueError( + "--situ-glu is mutually exclusive with --swiglu, --quick-geglu, " + "and --squared-relu." + ) + args.use_te_activation_func = True + args.bias_swiglu_fusion = False + if args.ffn_hidden_size is None: - if args.swiglu: + if args.swiglu or args.use_situ_glu: # reduce the dimnesion for MLP since projections happens on # two linear layers. this keeps the number of paramters in # the same ballpark as the counterpart with 4*h size @@ -2315,6 +2328,8 @@ def _add_network_size_args(parser): "bias_dropout_fusion", "apply_rope_fusion", "mamba_training_ssm_states_dtype", + # handled with the other model activation flags + "use_situ_glu", # internal/derived: controlled only via --tensor-parallel-num-weight-shards "gtp_weight_remat_size", # internal/derived: controlled only via --expert-tensor-parallel-num-weight-shards @@ -2399,6 +2414,16 @@ def _add_network_size_args(parser): help='Use squared relu activation instead of default gelu') group.add_argument('--swiglu', action='store_true', help='Use gated linear units and SiLU activation instead of default gelu') + group.add_argument( + '--situ-glu', + '--moe-use-situ-glu', + action='store_true', + dest='use_situ_glu', + help=( + 'Use SiTU-GLU in all gated dense and MoE FFNs. Native Transformer Engine ' + 'operators are preferred; otherwise MCore uses its local CuTe DSL fallback.' + ), + ) group.add_argument('--quick-geglu', action='store_true', help='Use quick geglu activation instead of default gelu') group.add_argument('--onnx-safe', type=bool, required=False, @@ -2791,8 +2816,7 @@ def _add_rl_args(parser): return parser def _add_training_args(parser): - from megatron.training.config import TrainingConfig - from megatron.training.config import ProfilingConfig + from megatron.training.config import ProfilingConfig, TrainingConfig prof_factory = ArgumentGroupFactory(ProfilingConfig) prof_group = prof_factory.build_group(parser, "profiling") @@ -3675,7 +3699,7 @@ def _add_kitchen_quantization_arguments(parser: argparse.ArgumentParser): If kitchen isn't available, nothing to do here, return unchanged parser """ try: - from megatron.core.extensions.kitchen import KitchenSpecProvider, HAVE_KITCHEN + from megatron.core.extensions.kitchen import HAVE_KITCHEN, KitchenSpecProvider except (ImportError, ModuleNotFoundError): HAVE_KITCHEN = False diff --git a/megatron/training/checkpointing.py b/megatron/training/checkpointing.py index 040cdac4c9b..e4e0c0daf7a 100644 --- a/megatron/training/checkpointing.py +++ b/megatron/training/checkpointing.py @@ -1697,7 +1697,7 @@ def _maybe_localize(value): def preprocess_fsdp_dtensor_state_dict(args, raw_state_dict, model): state_dict = raw_state_dict.copy() handle_fp8_extra_state_case(state_dict['model']) - if args.swiglu: + if args.swiglu or getattr(args, "use_situ_glu", False): if 'optimizer' in state_dict: model_state_dict, optimizer_state_dict = handle_swiglu_in_state_dict( model, state_dict['model'], state_dict['optimizer'] diff --git a/megatron/training/theoretical_memory_usage.py b/megatron/training/theoretical_memory_usage.py index ee398d3bf66..78f58fedb48 100644 --- a/megatron/training/theoretical_memory_usage.py +++ b/megatron/training/theoretical_memory_usage.py @@ -18,7 +18,7 @@ def compute_weight_and_optimizer_memory(args, verbose=False): args.num_query_groups = args.num_attention_heads # MoE. num_experts = 1 if args.num_experts is None else args.num_experts - gated_linear_multiplier = 3 / 2 if args.swiglu else 1 + gated_linear_multiplier = 3 / 2 if args.swiglu or getattr(args, "use_situ_glu", False) else 1 shared_expert_ffn_hidden_size = ( 0 diff --git a/megatron/training/training.py b/megatron/training/training.py index 37ee411e019..62ceaa78515 100644 --- a/megatron/training/training.py +++ b/megatron/training/training.py @@ -1067,7 +1067,7 @@ def transformer_flops(): fma_expansion_factor = 2 # - 3x (SwiGLU enabled): h->2*ffn_h GEMM and ffn_h->h GEMM are stacked. # - 2x (SwiGLU disabled): h->ffn_h GEMM and ffn_h->h GEMM are stacked. - ffn_expansion_factor = 3 if args.swiglu else 2 + ffn_expansion_factor = 3 if args.swiglu or getattr(args, "use_situ_glu", False) else 2 # self_attn is split into a token-linear part (projections, multiplied by # ``batch_size * args.seq_length`` like all other token-linear work) and a @@ -1366,7 +1366,7 @@ def _split_spec_part(part): gqa_groups=args.num_query_groups, kv_channels=args.kv_channels, mlp_expansion=args.ffn_hidden_size / args.hidden_size, - swiglu=args.swiglu, + swiglu=args.swiglu or getattr(args, "use_situ_glu", False), use_gated_delta_product=_uses_gated_delta_product_spec(args), moe_latent_size=args.moe_latent_size, moe_ffn_hidden_size=(args.moe_ffn_hidden_size if args.moe_ffn_hidden_size is not None diff --git a/tests/unit_tests/fusions/test_cutedsl_situ_glu.py b/tests/unit_tests/fusions/test_cutedsl_situ_glu.py new file mode 100644 index 00000000000..f9d77ef6658 --- /dev/null +++ b/tests/unit_tests/fusions/test_cutedsl_situ_glu.py @@ -0,0 +1,332 @@ +# Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. + +import argparse +import os +from contextlib import nullcontext + +import pytest +import torch +import torch.nn.functional as F + +from megatron.core.extensions.transformer_engine import TEActivationOp +from megatron.core.fusions.cutedsl_situ_glu import ( + CuTeDSLScaledSiTUGLU, + CuTeDSLSiTUGLU, + make_scaled_situ_glu, + make_situ_glu, + situ_glu_reference, +) +from megatron.core.transformer.transformer_config import TransformerConfig + + +def test_situ_glu_reference_formula(): + input_ = torch.tensor([[0.5, -1.0, 2.0, -3.0]], dtype=torch.float64) + gate, up = input_.chunk(2, dim=-1) + expected = 4.0 * torch.tanh(gate / 4.0) * torch.sigmoid(gate) + expected = expected * 25.0 * torch.tanh(up / 25.0) + torch.testing.assert_close(situ_glu_reference(input_), expected) + + +@pytest.mark.parametrize("flag", ["--situ-glu", "--moe-use-situ-glu"]) +def test_situ_glu_cli_aliases_enable_the_global_activation(flag): + from megatron.training.arguments import _add_network_size_args + + parser = argparse.ArgumentParser() + _add_network_size_args(parser) + + args = parser.parse_args([flag]) + + assert args.use_situ_glu is True + + +@pytest.mark.parametrize( + "precision_kwargs", + [ + {"bf16": True}, + {"fp8": "hybrid", "fp8_recipe": "mxfp8"}, + {"fp4": "e2m1", "fp4_recipe": "nvfp4"}, + ], + ids=["bf16", "mxfp8", "nvfp4"], +) +def test_situ_glu_selects_common_te_activation_builder(monkeypatch, precision_kwargs): + import transformer_engine.pytorch as te + + monkeypatch.delattr(te.ops, "SiTUGLU", raising=False) + monkeypatch.delattr(te.ops, "ScaledSiTUGLU", raising=False) + config = TransformerConfig( + num_layers=1, + hidden_size=32, + num_attention_heads=4, + activation_func=F.silu, + gated_linear_unit=True, + use_te_activation_func=True, + use_situ_glu=True, + **precision_kwargs, + ) + + activation = TEActivationOp(config) + + assert isinstance(activation, CuTeDSLSiTUGLU) + + +def test_situ_glu_prefers_complete_native_te_interface(monkeypatch): + import transformer_engine.pytorch as te + + class NativeSiTUGLU(torch.nn.Module): + def __init__(self, *, beta1, beta2, cache_quantized_input=False, glu_interleave_size=None): + super().__init__() + self.beta1 = beta1 + self.beta2 = beta2 + self.cache_quantized_input = cache_quantized_input + self.glu_interleave_size = glu_interleave_size + + class NativeScaledSiTUGLU(torch.nn.Module): + def __init__( + self, glu_interleave_size=None, *, activation_recompute_in_mlp=False, beta1, beta2 + ): + super().__init__() + self.beta1 = beta1 + self.beta2 = beta2 + self.glu_interleave_size = glu_interleave_size + self.activation_recompute_in_mlp = activation_recompute_in_mlp + + monkeypatch.setattr(te.ops, "SiTUGLU", NativeSiTUGLU, raising=False) + monkeypatch.setattr(te.ops, "ScaledSiTUGLU", NativeScaledSiTUGLU, raising=False) + + activation = make_situ_glu(beta1=4.0, beta2=25.0) + scaled_activation = make_scaled_situ_glu(beta1=4.0, beta2=25.0, glu_interleave_size=32) + + assert isinstance(activation, NativeSiTUGLU) + assert activation.beta1 == 4.0 + assert activation.beta2 == 25.0 + assert isinstance(scaled_activation, NativeScaledSiTUGLU) + assert scaled_activation.glu_interleave_size == 32 + assert not scaled_activation.activation_recompute_in_mlp + + +def test_situ_glu_falls_back_when_te_interface_is_incomplete(monkeypatch): + import transformer_engine.pytorch as te + + monkeypatch.setattr(te.ops, "SiTUGLU", torch.nn.Identity, raising=False) + monkeypatch.delattr(te.ops, "ScaledSiTUGLU", raising=False) + + assert isinstance(make_situ_glu(), CuTeDSLSiTUGLU) + assert isinstance(make_scaled_situ_glu(), CuTeDSLScaledSiTUGLU) + + +def test_swiglu_keeps_existing_te_activation_path(): + import transformer_engine.pytorch as te + + config = TransformerConfig( + num_layers=1, + hidden_size=32, + num_attention_heads=4, + activation_func=F.silu, + gated_linear_unit=True, + use_te_activation_func=True, + use_situ_glu=False, + ) + + assert isinstance(TEActivationOp(config), te.ops.SwiGLU) + + +def test_scaled_situ_glu_bf16_fallback_does_not_require_cudnn(monkeypatch): + import transformer_engine.pytorch as te + + monkeypatch.delattr(te.ops, "SiTUGLU", raising=False) + monkeypatch.delattr(te.ops, "ScaledSiTUGLU", raising=False) + activation = make_scaled_situ_glu(beta1=4.0, beta2=25.0, glu_interleave_size=None) + + assert isinstance(activation, CuTeDSLScaledSiTUGLU) + assert not isinstance(activation, te.ops.ScaledSwiGLU) + + +def test_scaled_situ_glu_fallback_rejects_activation_recompute(monkeypatch): + import transformer_engine.pytorch as te + + monkeypatch.delattr(te.ops, "SiTUGLU", raising=False) + monkeypatch.delattr(te.ops, "ScaledSiTUGLU", raising=False) + + with pytest.raises(ValueError, match="does not support activation recomputation"): + make_scaled_situ_glu(activation_recompute_in_mlp=True) + + +@pytest.mark.parametrize( + ("precision_kwargs", "error"), + [ + ({"fp8": "hybrid", "fp8_recipe": "delayed"}, "supports FP8 only"), + ( + {"fp4": "e2m1", "fp4_recipe": "custom", "fp4_quantizer_factory": "test.fake.factory"}, + "supports FP4 only", + ), + ], +) +def test_situ_glu_rejects_unsupported_quantization_recipe(precision_kwargs, error): + with pytest.raises(ValueError, match=error): + TransformerConfig( + num_layers=1, + hidden_size=32, + num_attention_heads=4, + activation_func=F.silu, + gated_linear_unit=True, + use_te_activation_func=True, + use_situ_glu=True, + **precision_kwargs, + ) + + +@pytest.mark.parametrize("dtype", [torch.bfloat16, torch.float16]) +def test_cutedsl_situ_glu_forward_backward(dtype): + torch.manual_seed(1234) + input_ref = torch.randn(5, 64, device="cuda", dtype=dtype, requires_grad=True) + input_cute = input_ref.detach().clone().requires_grad_(True) + grad = torch.randn(5, 32, device="cuda", dtype=dtype) + + output_ref = situ_glu_reference(input_ref) + output_ref.backward(grad) + + output_cute = CuTeDSLSiTUGLU()(input_cute) + output_cute.backward(grad) + + torch.testing.assert_close(output_cute, output_ref, rtol=2.0e-2, atol=2.0e-2) + torch.testing.assert_close(input_cute.grad, input_ref.grad, rtol=3.0e-2, atol=3.0e-2) + + +@pytest.mark.parametrize("dtype", [torch.bfloat16, torch.float16]) +def test_cutedsl_situ_glu_interleaved_forward_backward(dtype): + torch.manual_seed(4321) + rows = 5 + width = 64 + interleave_size = 8 + input_ref = torch.randn(rows, width, device="cuda", dtype=dtype, requires_grad=True) + input_interleaved = ( + input_ref.detach() + .reshape(rows, 2, width // (2 * interleave_size), interleave_size) + .transpose(1, 2) + .contiguous() + .view(rows, width) + .requires_grad_(True) + ) + grad = torch.randn(rows, width // 2, device="cuda", dtype=dtype) + + output_ref = situ_glu_reference(input_ref) + output_ref.backward(grad) + + output_cute = CuTeDSLSiTUGLU(interleave_size=interleave_size)(input_interleaved) + output_cute.backward(grad) + grad_interleaved_ref = ( + input_ref.grad.reshape(rows, 2, width // (2 * interleave_size), interleave_size) + .transpose(1, 2) + .contiguous() + .view(rows, width) + ) + + torch.testing.assert_close(output_cute, output_ref, rtol=2.0e-2, atol=2.0e-2) + torch.testing.assert_close( + input_interleaved.grad, grad_interleaved_ref, rtol=3.0e-2, atol=3.0e-2 + ) + + +@pytest.mark.parametrize("activation", ["swiglu", "situglu"]) +@pytest.mark.parametrize("precision", ["bf16", "mxfp8", "nvfp4"]) +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA not available") +@pytest.mark.internal +def test_mcore_moe_glu_forward_backward_uses_expected_backend(precision, activation): + """Run a real MCore MoE and prove SwiGLU and SiTU-GLU select the expected backend.""" + if torch.cuda.get_device_capability()[0] < 10: + pytest.skip("Fused block-scaled grouped MLP requires Blackwell") + if int(os.environ.get("NVTE_CUTEDSL_FUSED_GROUPED_MLP", "0")) <= 0: + pytest.skip("NVTE_CUTEDSL_FUSED_GROUPED_MLP is not enabled") + + import transformer_engine.pytorch as te + + from megatron.core.fp4_utils import get_fp4_context + from megatron.core.fp8_utils import get_fp8_context + from megatron.core.models.gpt.gpt_layer_specs import ( + get_gpt_layer_with_transformer_engine_submodules, + ) + from megatron.core.transformer.module import Float16Module + from megatron.core.transformer.moe.experts import TEGroupedMLP + from megatron.core.transformer.moe.moe_layer import MoELayer, MoESubmodules + from megatron.core.transformer.spec_utils import get_submodules + from megatron.training.initialize import _set_random_seed + from tests.unit_tests.test_utilities import Utils + + Utils.destroy_model_parallel() + Utils.initialize_model_parallel(1, 1) + precision_kwargs = {} + if precision == "mxfp8": + precision_kwargs = {"fp8": "hybrid", "fp8_recipe": "mxfp8"} + elif precision == "nvfp4": + precision_kwargs = {"fp4": "e2m1", "fp4_recipe": "nvfp4"} + + config = TransformerConfig( + num_layers=1, + hidden_size=256, + ffn_hidden_size=512, + num_attention_heads=8, + num_moe_experts=4, + moe_router_topk=2, + moe_router_load_balancing_type="none", + moe_token_dispatcher_type="alltoall", + moe_grouped_gemm=True, + use_transformer_engine_op_fuser=True, + moe_mlp_glu_interleave_size=32, + use_cpu_initialization=False, + add_bias_linear=False, + gated_linear_unit=True, + activation_func=F.silu, + use_te_activation_func=True, + use_situ_glu=activation == "situglu", + bias_activation_fusion=False, + bf16=True, + params_dtype=torch.bfloat16, + **precision_kwargs, + ) + _set_random_seed(seed_=123, data_parallel_random_init=False) + submodules = get_submodules( + get_gpt_layer_with_transformer_engine_submodules( + config.num_moe_experts, moe_grouped_gemm=True + ).mlp + ) + assert isinstance(submodules, MoESubmodules) + layer = MoELayer(config, submodules) + layer = Float16Module(layer.config, layer).module.cuda() + assert isinstance(layer.experts, TEGroupedMLP) + assert layer.experts._with_fused_impl + + hidden_states = torch.randn( + (4096, 1, config.hidden_size), dtype=torch.bfloat16, device="cuda", requires_grad=True + ) + context = nullcontext() + if precision == "mxfp8": + context = get_fp8_context(config) + elif precision == "nvfp4": + context = get_fp4_context(config) + with context: + output, _ = layer(hidden_states) + output.float().square().mean().backward() + + assert hidden_states.grad is not None + assert torch.isfinite(output).all() + assert all(parameter.grad is not None for parameter in layer.experts.parameters()) + (ops,) = layer.experts._fused_ops + if activation == "situglu": + native_scaled_situ_glu = getattr(te.ops, "ScaledSiTUGLU", None) + if native_scaled_situ_glu is None: + assert isinstance(ops[1], CuTeDSLScaledSiTUGLU) + else: + assert isinstance(ops[1], native_scaled_situ_glu) + expect_fused_glu = native_scaled_situ_glu is not None + else: + assert isinstance(ops[1], te.ops.ScaledSwiGLU) + expect_fused_glu = True + assert ops._module_groups is not None + fuser = ops._module_groups[0] + fused_types = tuple(type(op) for op, _ in fuser._forward_ops) + if expect_fused_glu and precision in ("mxfp8", "nvfp4"): + assert te.ops.fused.GroupedMLP_CuTeGEMMGLU in fused_types + else: + assert te.ops.fused.GroupedMLP_CuTeGEMMGLU not in fused_types + + Utils.destroy_model_parallel() diff --git a/tests/unit_tests/test_num_floating_point_operations.py b/tests/unit_tests/test_num_floating_point_operations.py index df5e4191843..3fcca2ca7af 100644 --- a/tests/unit_tests/test_num_floating_point_operations.py +++ b/tests/unit_tests/test_num_floating_point_operations.py @@ -100,6 +100,17 @@ def _make_hybrid_args(*, num_layers=4, hidden_size=512, num_attention_heads=8, s return args +def test_situ_glu_counts_the_same_ffn_gemms_as_swiglu(): + swiglu_args = _make_gpt_args(swiglu=True) + swiglu_args.use_situ_glu = False + situ_glu_args = _make_gpt_args(swiglu=False) + situ_glu_args.use_situ_glu = True + + assert num_floating_point_operations( + situ_glu_args, batch_size=8 + ) == num_floating_point_operations(swiglu_args, batch_size=8) + + class TestBSHDBackwardCompat: """For unpacked BSHD, the new optional arg must not change the result.""" diff --git a/tests/unit_tests/training/test_weight_and_optimizer_memory.py b/tests/unit_tests/training/test_weight_and_optimizer_memory.py index 19a532c35c2..7c0ee3a0ea4 100644 --- a/tests/unit_tests/training/test_weight_and_optimizer_memory.py +++ b/tests/unit_tests/training/test_weight_and_optimizer_memory.py @@ -33,6 +33,7 @@ def _make_args(**overrides): swiglu=False, tensor_model_parallel_size=2, untie_embeddings_and_output_weights=False, + use_situ_glu=False, use_distributed_optimizer=True, world_size=32, ) @@ -41,6 +42,15 @@ def _make_args(**overrides): return args +def test_situ_glu_counts_the_same_ffn_parameters_as_swiglu(): + swiglu_args = _make_args(swiglu=True) + situ_glu_args = _make_args(use_situ_glu=True) + + assert compute_weight_and_optimizer_memory( + situ_glu_args + ) == compute_weight_and_optimizer_memory(swiglu_args) + + def test_weight_and_optimizer_memory_accounts_for_expert_parallelism(): args = _make_args(pipeline_model_parallel_size=2, world_size=64) diff --git a/tests/unit_tests/transformer/moe/test_shared_experts.py b/tests/unit_tests/transformer/moe/test_shared_experts.py index c71f665259f..02da5eefa00 100644 --- a/tests/unit_tests/transformer/moe/test_shared_experts.py +++ b/tests/unit_tests/transformer/moe/test_shared_experts.py @@ -141,6 +141,9 @@ def _fake_shared_expert(**config_kwargs): fp4_recipe="nvfp4", fp8=True, fp8_recipe="mxfp8", + use_situ_glu=False, + situ_glu_beta1=4.0, + situ_glu_beta2=25.0, ) for key, value in config_kwargs.items(): setattr(config, key, value) @@ -218,6 +221,7 @@ def test_validate_fused_grouped_swiglu_requires_clamped_te_support(monkeypatch, ({"activation_func": F.gelu}, None, "SwiGLU activation"), ({"gated_linear_unit": False}, None, "SwiGLU activation"), ({"moe_shared_expert_glu_interleave_size": None}, None, "glu_interleave_size"), + ({"use_situ_glu": True, "moe_shared_expert_glu_interleave_size": 16}, None, "32-wide"), ({}, "linear_fc1", "FC1"), ({}, "linear_fc2", "FC2"), ], @@ -267,6 +271,31 @@ def test_make_fused_grouped_swiglu_ops_builds_grouped_pipeline(monkeypatch): assert fc2_op.weight0 is shared_expert.linear_fc2.weight +def test_make_fused_grouped_situ_glu_ops_uses_local_or_native_builder(monkeypatch): + from megatron.core.fusions import cutedsl_situ_glu + + _patch_fake_shared_expert_te(monkeypatch) + shared_expert = _fake_shared_expert(use_situ_glu=True) + calls = [] + + class FakeScaledSiTUGLU(torch.nn.Module): + pass + + def make_scaled_situ_glu(**kwargs): + calls.append(kwargs) + return FakeScaledSiTUGLU() + + monkeypatch.setattr(cutedsl_situ_glu, "make_scaled_situ_glu", make_scaled_situ_glu) + + ops = shared_expert._make_fused_grouped_swiglu_ops() + + _, activation_op, _ = list(ops.children()) + assert isinstance(activation_op, FakeScaledSiTUGLU) + assert calls == [ + {"beta1": 4.0, "beta2": 25.0, "install_grouped_fallback": True, "glu_interleave_size": 32} + ] + + def test_make_fused_grouped_swiglu_ops_builds_clamped_activation(monkeypatch): _patch_fake_shared_expert_te(monkeypatch) shared_expert = _fake_shared_expert(activation_func_clamp_value=7.0)