Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
6 changes: 5 additions & 1 deletion deepspec/modeling/dspark/gemma4/config.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
30 changes: 27 additions & 3 deletions deepspec/modeling/dspark/gemma4/modeling.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
)
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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
Expand Down