From e5c58eb9d627421236f78053137562b97c14984e Mon Sep 17 00:00:00 2001 From: Yuan Luo Date: Tue, 5 May 2026 00:36:26 +0800 Subject: [PATCH] [VLM] Optimize Gemma4 VLM with PCG and fuse RMSNorm + residual add + scalar (#24048) Co-authored-by: luoyuan.luo --- python/sglang/srt/layers/gemma4_fused_ops.py | 91 +++++++++++++++++++ .../sglang/srt/model_executor/model_runner.py | 14 ++- .../piecewise_cuda_graph_runner.py | 8 +- python/sglang/srt/models/gemma4_causal.py | 31 ++++++- python/sglang/srt/models/gemma4_mm.py | 20 ++++ 5 files changed, 158 insertions(+), 6 deletions(-) diff --git a/python/sglang/srt/layers/gemma4_fused_ops.py b/python/sglang/srt/layers/gemma4_fused_ops.py index 5f227db82..eeb9ebcea 100644 --- a/python/sglang/srt/layers/gemma4_fused_ops.py +++ b/python/sglang/srt/layers/gemma4_fused_ops.py @@ -77,3 +77,94 @@ def gemma_rmsnorm_residual_scalar( BLOCK_SIZE=BLOCK_SIZE, ) return out + + +@triton.jit +def _gemma_dual_rmsnorm_residual_kernel( + X1_ptr, + W1_ptr, + X2_ptr, + W2_ptr, + W3_ptr, + Residual_ptr, + Scalar_ptr, + Out_ptr, + stride_x1, + stride_x2, + stride_r, + stride_o, + N, + eps1, + eps2, + eps3, + BLOCK_SIZE: tl.constexpr, +): + """Fused: out = (rmsnorm(rmsnorm(x1,w1) + rmsnorm(x2,w2), w3) + residual) * scalar""" + row = tl.program_id(0) + cols = tl.arange(0, BLOCK_SIZE) + mask = cols < N + + x1 = tl.load(X1_ptr + row * stride_x1 + cols, mask=mask, other=0.0).to(tl.float32) + w1 = tl.load(W1_ptr + cols, mask=mask, other=0.0).to(tl.float32) + x2 = tl.load(X2_ptr + row * stride_x2 + cols, mask=mask, other=0.0).to(tl.float32) + w2 = tl.load(W2_ptr + cols, mask=mask, other=0.0).to(tl.float32) + w3 = tl.load(W3_ptr + cols, mask=mask, other=0.0).to(tl.float32) + r = tl.load(Residual_ptr + row * stride_r + cols, mask=mask, other=0.0).to( + tl.float32 + ) + + var1 = tl.sum(x1 * x1, axis=0) / N + norm1 = x1 * tl.rsqrt(var1 + eps1) * w1 + + var2 = tl.sum(x2 * x2, axis=0) / N + norm2 = x2 * tl.rsqrt(var2 + eps2) * w2 + + combined = norm1 + norm2 + + var3 = tl.sum(combined * combined, axis=0) / N + norm3 = combined * tl.rsqrt(var3 + eps3) * w3 + + scalar = tl.load(Scalar_ptr).to(tl.float32) + out = (norm3 + r) * scalar + + tl.store(Out_ptr + row * stride_o + cols, out.to(x1.dtype), mask=mask) + + +def gemma_dual_rmsnorm_residual_scalar( + x1: torch.Tensor, + weight1: torch.Tensor, + x2: torch.Tensor, + weight2: torch.Tensor, + weight3: torch.Tensor, + residual: torch.Tensor, + scalar: torch.Tensor, + eps1: float = 1e-6, + eps2: float = 1e-6, + eps3: float = 1e-6, +) -> torch.Tensor: + """Fused (rmsnorm(rmsnorm(x1,w1) + rmsnorm(x2,w2), w3) + residual) * scalar.""" + assert x1.dim() == 2 and x1.stride(-1) == 1 + M, N = x1.shape + BLOCK_SIZE = triton.next_power_of_2(N) + out = torch.empty_like(x1) + + _gemma_dual_rmsnorm_residual_kernel[(M,)]( + x1, + weight1, + x2, + weight2, + weight3, + residual, + scalar, + out, + x1.stride(0), + x2.stride(0), + residual.stride(0), + out.stride(0), + N, + eps1, + eps2, + eps3, + BLOCK_SIZE=BLOCK_SIZE, + ) + return out diff --git a/python/sglang/srt/model_executor/model_runner.py b/python/sglang/srt/model_executor/model_runner.py index 96d014a9e..e6f02f154 100644 --- a/python/sglang/srt/model_executor/model_runner.py +++ b/python/sglang/srt/model_executor/model_runner.py @@ -278,6 +278,10 @@ def resolve_language_model(model: nn.Module) -> nn.Module: model_cls_name = model.__class__.__name__ if model_cls_name == "Qwen3OmniMoeForConditionalGeneration": return model.thinker.model + if hasattr(model, "model"): + return model.model + if hasattr(model, "language_model"): + return model.language_model return model.model @@ -2834,8 +2838,12 @@ class ModelRunner(ModelRunnerKVCacheMixin): self.model.model = resolve_language_model(self.model) language_model = getattr(self.model, "language_model", self.model) - # Some draft models (e.g. eagle3) don't have a standard 'layers' attribute - if not hasattr(language_model.model, "layers"): + # Resolve model with layers: handle CausalLM wrapper (.model.layers) and direct TextModel (.layers) + if hasattr(language_model, "model") and hasattr(language_model.model, "layers"): + layer_model = language_model.model + elif hasattr(language_model, "layers"): + layer_model = language_model + else: logger.warning( "Disable piecewise CUDA graph because the model does not have a 'layers' attribute" ) @@ -2844,7 +2852,7 @@ class ModelRunner(ModelRunnerKVCacheMixin): self.attention_layers = [] self.moe_layers = [] self.moe_fusions = [] - for layer in language_model.model.layers: + for layer in layer_model.layers: attn_layer = None if hasattr(layer, "self_attn"): if hasattr(layer.self_attn, "attn"): diff --git a/python/sglang/srt/model_executor/piecewise_cuda_graph_runner.py b/python/sglang/srt/model_executor/piecewise_cuda_graph_runner.py index 178345292..c89fa5958 100644 --- a/python/sglang/srt/model_executor/piecewise_cuda_graph_runner.py +++ b/python/sglang/srt/model_executor/piecewise_cuda_graph_runner.py @@ -304,8 +304,14 @@ class PiecewiseCudaGraphRunner: language_model = getattr( self.model_runner.model, "language_model", self.model_runner.model ) + layer_model = ( + language_model.model + if hasattr(language_model, "model") + and hasattr(language_model.model, "layers") + else language_model + ) with patch_model( - language_model.model, self.compile_config.compiler + layer_model, self.compile_config.compiler ) as patched_model: # Dummy warmup for jit kernel diff --git a/python/sglang/srt/models/gemma4_causal.py b/python/sglang/srt/models/gemma4_causal.py index 544406119..1cf761d31 100644 --- a/python/sglang/srt/models/gemma4_causal.py +++ b/python/sglang/srt/models/gemma4_causal.py @@ -27,7 +27,10 @@ from transformers import ( from sglang.srt.distributed import ( get_tensor_model_parallel_world_size, ) -from sglang.srt.layers.gemma4_fused_ops import gemma_rmsnorm_residual_scalar +from sglang.srt.layers.gemma4_fused_ops import ( + gemma_dual_rmsnorm_residual_scalar, + gemma_rmsnorm_residual_scalar, +) from sglang.srt.layers.layernorm import Gemma4RMSNorm, RMSNorm from sglang.srt.layers.linear import ( QKVParallelLinear, @@ -545,12 +548,36 @@ class Gemma4DecoderLayer(nn.Module): # Dense MLP branch hidden_states_1 = self.mlp(hidden_states) - hidden_states_1 = self.post_feedforward_layernorm_1(hidden_states_1) # MoE branch: router sees residual (= post_attn_out + old_residual) router_logits = self.router(moe_input) hidden_states_2 = self.pre_feedforward_layernorm_2(moe_input) hidden_states_2 = self.moe(hidden_states_2, router_logits) + + # Fused: (rmsnorm(rmsnorm(h1,w1) + rmsnorm(h2,w2), w3) + residual) * scalar + if ( + not self.has_ple + and hidden_states_1.is_cuda + and hidden_states_1.dim() == 2 + ): + norm1 = self.post_feedforward_layernorm_1 + norm2 = self.post_feedforward_layernorm_2 + norm3 = self.post_feedforward_layernorm + hidden_states = gemma_dual_rmsnorm_residual_scalar( + hidden_states_1, + norm1.weight.data, + hidden_states_2, + norm2.weight.data, + norm3.weight.data, + residual, + self.layer_scalar, + norm1.variance_epsilon, + norm2.variance_epsilon, + norm3.variance_epsilon, + ) + return hidden_states, None + + hidden_states_1 = self.post_feedforward_layernorm_1(hidden_states_1) hidden_states_2 = self.post_feedforward_layernorm_2(hidden_states_2) # Combine branches diff --git a/python/sglang/srt/models/gemma4_mm.py b/python/sglang/srt/models/gemma4_mm.py index 4618129fa..a1f0fe2fe 100644 --- a/python/sglang/srt/models/gemma4_mm.py +++ b/python/sglang/srt/models/gemma4_mm.py @@ -224,6 +224,26 @@ class Gemma4ForConditionalGeneration(PreTrainedModel): self.post_init() + @property + def model(self): + # Alias .model to .language_model so this class satisfies the piecewise + # CUDA graph gate (which checks `hasattr(model, "model")`). Implemented + # as a property to avoid registering a duplicate submodule in + # `_modules`, which would double state_dict keys and disturb + # ShardedStateLoader / CPU-offload / dummy-init paths. + return self.language_model + + def __setattr__(self, name, value): + # Block writes to "model" so the runner's + # `self.model.model = resolve_language_model(self.model)` (which for + # this class returns language_model itself) is a no-op rather than a + # nn.Module submodule registration. Without this, nn.Module.__setattr__ + # would bypass the @property's setter for Module values and pollute + # `_modules` with a duplicate alias, doubling state_dict keys. + if name == "model": + return + super().__setattr__(name, value) + def pad_input_ids( self, input_ids: List[int],