From 8277cecf90da3eaeae87fffda41809175efd269b Mon Sep 17 00:00:00 2001 From: GitHub Contributor Bot Date: Sat, 11 Jul 2026 18:43:48 +0530 Subject: [PATCH] fix: Fix: Feat: Compatible with Gamme4 MoE structure --- deepspec/modeling/dspark/gemma4/config.py | 6 ++++- deepspec/modeling/dspark/gemma4/modeling.py | 30 ++++++++++++++++++--- 2 files changed, 32 insertions(+), 4 deletions(-) diff --git a/deepspec/modeling/dspark/gemma4/config.py b/deepspec/modeling/dspark/gemma4/config.py index 3c10f166..48efbb77 100644 --- a/deepspec/modeling/dspark/gemma4/config.py +++ b/deepspec/modeling/dspark/gemma4/config.py @@ -7,8 +7,12 @@ def get_gemma4_text_config(target_config): + if target_config.model_type in ("gemma4_text", "gemma4_unified_text"): + return copy.deepcopy(target_config) + assert target_config.model_type in ("gemma4", "gemma4_unified"), ( - "Gemma4 DSpark expects a Gemma4 or Gemma4 Unified top-level target config, " + "Gemma4 DSpark expects a Gemma4/Gemma4 Unified top-level target config " + "or a Gemma4/Gemma4 Unified text config, " f"got model_type={target_config.model_type!r}." ) text_config = target_config.text_config diff --git a/deepspec/modeling/dspark/gemma4/modeling.py b/deepspec/modeling/dspark/gemma4/modeling.py index c8d70c31..ab70bd58 100644 --- a/deepspec/modeling/dspark/gemma4/modeling.py +++ b/deepspec/modeling/dspark/gemma4/modeling.py @@ -11,8 +11,10 @@ from transformers.models.gemma4.modeling_gemma4 import ( Gemma4PreTrainedModel, Gemma4RMSNorm, + Gemma4TextExperts, Gemma4TextMLP, Gemma4TextRotaryEmbedding, + Gemma4TextRouter, Gemma4TextScaledWordEmbedding, apply_rotary_pos_emb as apply_gemma4_rotary_pos_emb, ) @@ -172,14 +174,27 @@ class Gemma4DSparkDecoderLayer(GradientCheckpointingLayer): def __init__(self, config, layer_idx: int): super().__init__() self.hidden_size = config.hidden_size - assert not bool(config.enable_moe_block), ( - "Gemma4 DSpark prototype does not support Gemma4 MoE blocks yet." - ) + self.enable_moe_block = bool(config.enable_moe_block) assert int(config.hidden_size_per_layer_input) == 0, ( "Gemma4 DSpark prototype does not support per-layer input gates yet." ) self.self_attn = Gemma4DSparkAttention(config=config, layer_idx=layer_idx) self.mlp = Gemma4TextMLP(config, layer_idx) + if self.enable_moe_block: + self.router = Gemma4TextRouter(config) + self.experts = Gemma4TextExperts(config) + self.post_feedforward_layernorm_1 = Gemma4RMSNorm( + config.hidden_size, + eps=config.rms_norm_eps, + ) + self.post_feedforward_layernorm_2 = Gemma4RMSNorm( + config.hidden_size, + eps=config.rms_norm_eps, + ) + self.pre_feedforward_layernorm_2 = Gemma4RMSNorm( + config.hidden_size, + eps=config.rms_norm_eps, + ) self.input_layernorm = Gemma4RMSNorm( config.hidden_size, eps=config.rms_norm_eps, @@ -233,6 +248,15 @@ def forward( residual = hidden_states hidden_states = self.pre_feedforward_layernorm(hidden_states) hidden_states = self.mlp(hidden_states) + if self.enable_moe_block: + hidden_states_1 = self.post_feedforward_layernorm_1(hidden_states) + hidden_states_flat = residual.reshape(-1, residual.shape[-1]) + _, top_k_weights, top_k_index = self.router(hidden_states_flat) + hidden_states_2 = self.pre_feedforward_layernorm_2(hidden_states_flat) + hidden_states_2 = self.experts(hidden_states_2, top_k_index, top_k_weights) + hidden_states_2 = hidden_states_2.reshape(residual.shape) + hidden_states_2 = self.post_feedforward_layernorm_2(hidden_states_2) + hidden_states = hidden_states_1 + hidden_states_2 hidden_states = self.post_feedforward_layernorm(hidden_states) hidden_states = residual + hidden_states return hidden_states * self.layer_scalar