diff --git a/src/mobius/integrations/ort_genai/auto_export_test.py b/src/mobius/integrations/ort_genai/auto_export_test.py index ac9aa1ff..0b0c60f8 100644 --- a/src/mobius/integrations/ort_genai/auto_export_test.py +++ b/src/mobius/integrations/ort_genai/auto_export_test.py @@ -1055,7 +1055,7 @@ class FakeConfig: decoder = _mock_model_with_inputs( [ "inputs_embeds", - "input_ids", + "per_layer_inputs", "attention_mask", "position_ids", "past_key_values.0.key", @@ -1108,8 +1108,8 @@ def test_gemma4_vision_inputs(self, tmp_path): assert "image_grid_thw" not in vision_inputs assert data["model"]["vision"]["spatial_merge_size"] == 2 - def test_gemma4_decoder_has_input_ids_and_inputs_embeds(self, tmp_path): - """Gemma4 decoder has both inputs_embeds and input_ids.""" + def test_gemma4_decoder_has_per_layer_inputs_and_inputs_embeds(self, tmp_path): + """Gemma4 decoder has inputs_embeds and per_layer_inputs.""" pkg = self._make_gemma4_pkg() path = _write_genai_config( pkg.config, @@ -1128,7 +1128,8 @@ def test_gemma4_decoder_has_input_ids_and_inputs_embeds(self, tmp_path): data = json.load(f) decoder_inputs = data["model"]["decoder"]["inputs"] assert "inputs_embeds" in decoder_inputs - assert "input_ids" in decoder_inputs + assert "per_layer_inputs" in decoder_inputs + assert "input_ids" not in decoder_inputs # KV cache templates are present assert decoder_inputs["past_key_names"] == "past_key_values.%d.key" @@ -1342,7 +1343,7 @@ def test_gemma4_genai_config_from_real_model(self, tmp_path): # Decoder inputs introspected from graph decoder_inputs = data["model"]["decoder"]["inputs"] assert "inputs_embeds" in decoder_inputs - assert "input_ids" in decoder_inputs + assert "input_ids" not in decoder_inputs assert "attention_mask" in decoder_inputs assert "position_ids" in decoder_inputs assert decoder_inputs["past_key_names"] == ("past_key_values.%d.key") diff --git a/src/mobius/models/gemma4.py b/src/mobius/models/gemma4.py index 57ae1c4b..1ec109dc 100644 --- a/src/mobius/models/gemma4.py +++ b/src/mobius/models/gemma4.py @@ -1383,13 +1383,12 @@ def __init__(self, config: Gemma4Config): self.rotary_emb_local = initialize_rope(local_config) self.rotary_emb_global = initialize_rope(global_config) - # Per-layer input embeddings (optional feature) + # Per-layer input dimension (used by decoder layers). + # In VLM 3-model split, per-layer inputs are precomputed by the + # embedding model and passed as per_layer_inputs. In single-model + # (text-only) mode, they are computed here from input_ids. self._per_layer_dim = getattr(config, "hidden_size_per_layer_input", 0) self._hidden_size = config.hidden_size - # Multimodal token IDs to mask before per-layer embedding lookup. - # HF masks these to pad_token_id (0) so image/audio slots don't contribute - # arbitrary large-ID embeddings to the per-layer gate (see - # Gemma4Model.forward lines 39-46 in HuggingFace transformers). self._image_token_id: int = config.image_token_id or 0 self._audio_token_id: int | None = ( config.audio.audio_token_id if config.audio is not None else None @@ -1397,20 +1396,13 @@ def __init__(self, config: Gemma4Config): if self._per_layer_dim: self._num_layers = config.num_hidden_layers vocab_per_layer = getattr(config, "vocab_size_per_layer_input", 0) - # Use per-layer embedding tables instead of one giant [V, L*D] table. - # Each [V, D] table has only V*D elements (e.g. 262144*256 = 67M), - # well under the ORT CUDA Gather int32 limit (~2.1B). This also - # avoids the post-Gather reshape and per-layer axis-2 slicing. - self.embed_tokens_per_layer = nn.ModuleList( - [ - Gemma3TextScaledWordEmbedding( - vocab_per_layer, - self._per_layer_dim, - config.pad_token_id, - embed_scale=float(self._per_layer_dim**0.5), - ) - for _ in range(config.num_hidden_layers) - ] + # Single fused [V, L*D] table. Requires ORT >= 1.27 for CUDA + # Gather int64 index support (onnxruntime#28107). + self.embed_tokens_per_layer = Gemma3TextScaledWordEmbedding( + vocab_per_layer, + self._num_layers * self._per_layer_dim, + config.pad_token_id, + embed_scale=float(self._per_layer_dim**0.5), ) self.per_layer_model_projection = Linear( config.hidden_size, @@ -1424,22 +1416,10 @@ def __init__(self, config: Gemma4Config): def _compute_per_layer_inputs( self, op: OpBuilder, - input_ids: ir.Value | None, + input_ids: ir.Value, inputs_embeds: ir.Value, - ) -> list[ir.Value] | None: - """Compute per-layer input embeddings, one ``[B, S, per_layer_dim]`` per layer. - - HF's ``Gemma4Model.forward`` replaces image/audio token positions with - ``pad_token_id`` (0) *before* calling ``embed_tokens_per_layer``. We - replicate that masking here so each layer's per-layer gate sees a PAD - embedding at multimodal positions rather than the raw soft-token IDs - (258880 / 258881), which are semantically meaningless in the text - per-layer vocabulary. - """ - if not self._per_layer_dim: - return None - - # Project hidden states and scale by hidden_size**-0.5 (matches HF) + ) -> list[ir.Value]: + """Compute per-layer input embeddings for single-model (text-only) mode.""" proj = self.per_layer_model_projection(op, inputs_embeds) proj = op.Mul(proj, float(self._hidden_size**-0.5)) proj = op.Reshape( @@ -1447,38 +1427,34 @@ def _compute_per_layer_inputs( ) proj = self.per_layer_projection_norm(op, proj) - # Mask multimodal token IDs to pad_token_id (0) - masked_ids: ir.Value | None = None - if input_ids is not None: - pad = op.Constant(value_int=0) - masked_ids = input_ids - if self._image_token_id: - masked_ids = op.Where( - op.Equal(masked_ids, op.Constant(value_int=self._image_token_id)), - pad, - masked_ids, - ) - if self._audio_token_id is not None: - masked_ids = op.Where( - op.Equal(masked_ids, op.Constant(value_int=self._audio_token_id)), - pad, - masked_ids, - ) - - # Per-layer embeddings: each table is [V, per_layer_dim] — small enough - # to avoid ORT CUDA Gather int32 overflow (onnxruntime#28107). - per_layer_results: list[ir.Value] = [] - for i in range(self._num_layers): - # Slice proj along axis 2 for this layer: [B, S, L, D] → [B, S, D] - proj_i = op.Squeeze(op.Slice(proj, starts=[i], ends=[i + 1], axes=[2]), [2]) + pad = op.Constant(value_int=0) + masked_ids = input_ids + if self._image_token_id: + masked_ids = op.Where( + op.Equal(masked_ids, op.Constant(value_int=self._image_token_id)), + pad, + masked_ids, + ) + if self._audio_token_id is not None: + masked_ids = op.Where( + op.Equal(masked_ids, op.Constant(value_int=self._audio_token_id)), + pad, + masked_ids, + ) - if masked_ids is not None: - token_emb_i = self.embed_tokens_per_layer[i](op, masked_ids) - proj_i = op.Add(proj_i, token_emb_i) + fused_emb = self.embed_tokens_per_layer(op, masked_ids) + fused_emb = op.Reshape( + fused_emb, + op.Constant(value_ints=[0, 0, self._num_layers, self._per_layer_dim]), + ) - per_layer_results.append(op.Mul(proj_i, float(0.5**0.5))) + combined = op.Add(proj, fused_emb) + combined = op.Mul(combined, float(0.5**0.5)) - return per_layer_results + return [ + op.Squeeze(op.Slice(combined, starts=[i], ends=[i + 1], axes=[2]), [2]) + for i in range(self._num_layers) + ] def forward( self, @@ -1488,13 +1464,30 @@ def forward( position_ids: ir.Value, past_key_values: list | None = None, inputs_embeds: ir.Value | None = None, + per_layer_inputs: ir.Value | None = None, ) -> tuple[ir.Value, list]: if inputs_embeds is not None: hidden_states = inputs_embeds else: hidden_states = self.embed_tokens(op, input_ids) - per_layer_inputs = self._compute_per_layer_inputs(op, input_ids, hidden_states) + # Unpack precomputed per_layer_inputs [B, S, L*D] (VLM split), + # or compute from input_ids (text-only single-model). + per_layer_list: list[ir.Value] | None = None + if self._per_layer_dim and per_layer_inputs is not None: + # VLM split: unpack precomputed per-layer inputs + num_layers = len(self.layers) + per_layer_4d = op.Reshape( + per_layer_inputs, + op.Constant(value_ints=[0, 0, num_layers, self._per_layer_dim]), + ) + per_layer_list = [ + op.Squeeze(op.Slice(per_layer_4d, starts=[i], ends=[i + 1], axes=[2]), [2]) + for i in range(num_layers) + ] + elif self._per_layer_dim and input_ids is not None: + # Text-only: compute per-layer inputs from input_ids + per_layer_list = self._compute_per_layer_inputs(op, input_ids, hidden_states) # Determine whether to emit GroupQueryAttention directly. # GQA fuses RoPE + attention + KV cache into a single op, and @@ -1606,7 +1599,7 @@ def forward( for i, (layer, layer_type, past_kv) in enumerate( zip(self.layers, self.layer_types, past_kvs) ): - per_layer_input = per_layer_inputs[i] if per_layer_inputs is not None else None + per_layer_input = per_layer_list[i] if per_layer_list is not None else None # Per-layer decision: use GQA when available. KV-shared layers # also use GQA (with empty K/V and shared past buffer). @@ -1700,17 +1693,8 @@ def preprocess_weights( state_dict[new_key] = state_dict.pop(key) elif "vision_tower" in key or "embed_vision" in key: state_dict.pop(key, None) - # Split fused per-layer embedding: HF stores one [V, L*D] tensor; - # we use nn.ModuleList of L separate [V, D] Embedding tables. - per_layer_dim = self.config.hidden_size_per_layer_input - if per_layer_dim > 0: - fused_key = "model.embed_tokens_per_layer.weight" - if fused_key in state_dict: - value = state_dict.pop(fused_key) - num_layers = self.config.num_hidden_layers - for i in range(num_layers): - shard = value[:, i * per_layer_dim : (i + 1) * per_layer_dim] - state_dict[f"model.embed_tokens_per_layer.{i}.weight"] = shard + # HF's model.embed_tokens_per_layer.weight [V, L*D] maps directly + # to our fused embedding table — no splitting needed. # Map HF expert weight names and fold router scale _remap_moe_expert_weights(state_dict, self.config) return super().preprocess_weights(state_dict) @@ -1724,12 +1708,10 @@ def preprocess_weights( class _Gemma4DecoderModel(nn.Module): """Gemma4 text decoder sub-model accepting ``inputs_embeds``. - When ``hidden_size_per_layer_input > 0`` (e.g. E2B), the text model also - needs the original ``input_ids`` to compute per-layer token embeddings that - condition each decoder layer. HF's ``Gemma4ForConditionalGeneration`` - passes *both* ``inputs_embeds`` and ``input_ids`` to the language model for - this reason. We mirror that by accepting ``input_ids`` as an optional - ONNX graph input and forwarding it to :class:`Gemma4TextModel`. + When ``hidden_size_per_layer_input > 0`` (e.g. Gemma4 E2B), per-layer input + embeddings are precomputed by the embedding sub-model and passed as + ``per_layer_inputs`` (shape ``[B, S, L*D]``). The decoder unpacks them + and feeds one ``[B, S, D]`` slice to each decoder layer's gating mechanism. """ def __init__(self, config: Gemma4Config): @@ -1744,16 +1726,17 @@ def forward( inputs_embeds: ir.Value, attention_mask: ir.Value, position_ids: ir.Value, - input_ids: ir.Value | None = None, + per_layer_inputs: ir.Value | None = None, past_key_values: list | None = None, ) -> tuple[ir.Value, list]: hidden_states, present_key_values = self.model( op, - input_ids=input_ids, + input_ids=None, attention_mask=attention_mask, position_ids=position_ids, past_key_values=past_key_values, inputs_embeds=inputs_embeds, + per_layer_inputs=per_layer_inputs, ) logits = self.lm_head(op, hidden_states) # Gemma4 applies final logit soft-capping: logit_cap * tanh(x / logit_cap) @@ -1853,9 +1836,9 @@ class Gemma4EmbeddingModel(nn.Module): has audio support (``config.audio is not None``), ``forward()`` also accepts ``audio_features`` and scatters them at audio-token positions. - Both image and audio use a one-row dummy guard so that ORT's eager - evaluation of both ``Where`` branches never ``Gather`` on a zero-length - tensor during text-only / decode steps. + When ``hidden_size_per_layer_input > 0``, also computes per-layer input + embeddings that condition each decoder layer. This moves the per-layer + computation out of the decoder so the decoder no longer needs ``input_ids``. Inputs (image-only variant): - ``input_ids [B, S]`` INT64 @@ -1866,7 +1849,9 @@ class Gemma4EmbeddingModel(nn.Module): - ``image_features [num_img_tokens, hidden_size]`` - ``audio_features [num_aud_tokens, hidden_size]`` - Output: ``inputs_embeds [B, S, hidden_size]`` + Outputs: + - ``inputs_embeds [B, S, hidden_size]`` + - ``per_layer_inputs [B, S, L*D]`` (only when ``hidden_size_per_layer_input > 0``) """ def __init__(self, config: Gemma4Config): @@ -1883,6 +1868,34 @@ def __init__(self, config: Gemma4Config): # Audio token ID is only set when the model has an audio encoder. self.audio_token_id: int | None = config.audio.audio_token_id if config.audio else None + # Per-layer input embedding components (moved from the decoder). + self._per_layer_dim = getattr(config, "hidden_size_per_layer_input", 0) + self._hidden_size = config.hidden_size + if self._per_layer_dim: + self._num_layers = config.num_hidden_layers + vocab_per_layer = getattr(config, "vocab_size_per_layer_input", 0) + # Single fused [V, L*D] embedding table matching HuggingFace's + # ``embed_tokens_per_layer.weight`` shape. The ORT CUDA Gather + # int32 overflow (onnxruntime#28107) is now fixed. + self.embed_tokens_per_layer = Gemma3TextScaledWordEmbedding( + vocab_per_layer, + self._num_layers * self._per_layer_dim, + config.pad_token_id, + embed_scale=float(self._per_layer_dim**0.5), + ) + self.per_layer_model_projection = Linear( + config.hidden_size, + config.num_hidden_layers * self._per_layer_dim, + bias=False, + ) + self.per_layer_projection_norm = RMSNorm( + self._per_layer_dim, eps=config.rms_norm_eps + ) + self._image_token_id_mask: int = config.image_token_id or 0 + self._audio_token_id_mask: int | None = ( + config.audio.audio_token_id if config.audio is not None else None + ) + def _scatter_features( self, op: OpBuilder, @@ -1925,7 +1938,7 @@ def forward( input_ids: ir.Value, image_features: ir.Value, audio_features: ir.Value | None = None, - ) -> ir.Value: + ) -> ir.Value | tuple[ir.Value, ir.Value]: # [B, S] → [B, S, hidden] hidden = self.embed_tokens(op, input_ids) @@ -1944,7 +1957,53 @@ def forward( op, hidden, input_ids, self.audio_token_id, audio_features ) - return hidden + if not self._per_layer_dim: + return hidden + + # Compute per-layer input embeddings (moved from the decoder). + # 1. Project hidden states → [B, S, L*D] and scale by hidden_size**-0.5 + proj = self.per_layer_model_projection(op, hidden) + proj = op.Mul(proj, float(self._hidden_size**-0.5)) + # Reshape to [B, S, L, D] for per-layer RMSNorm + proj = op.Reshape( + proj, op.Constant(value_ints=[0, 0, self._num_layers, self._per_layer_dim]) + ) + proj = self.per_layer_projection_norm(op, proj) + + # 2. Mask multimodal token IDs → pad_token_id (0) before per-layer lookup + pad = op.Constant(value_int=0) + masked_ids = input_ids + if self._image_token_id_mask: + masked_ids = op.Where( + op.Equal(masked_ids, op.Constant(value_int=self._image_token_id_mask)), + pad, + masked_ids, + ) + if self._audio_token_id_mask is not None: + masked_ids = op.Where( + op.Equal(masked_ids, op.Constant(value_int=self._audio_token_id_mask)), + pad, + masked_ids, + ) + + # 3. Single Gather on fused [V, L*D] table → reshape to [B, S, L, D] + fused_emb = self.embed_tokens_per_layer(op, masked_ids) + # fused_emb: [B, S, L*D] → [B, S, L, D] + fused_emb = op.Reshape( + fused_emb, + op.Constant(value_ints=[0, 0, self._num_layers, self._per_layer_dim]), + ) + + # 4. Combine: (proj + emb) * 0.707 per layer, then flatten back + combined = op.Add(proj, fused_emb) # [B, S, L, D] + combined = op.Mul(combined, float(0.5**0.5)) + # Flatten L*D → single per_layer_inputs output: [B, S, L*D] + per_layer_inputs = op.Reshape( + combined, + op.Constant(value_ints=[0, 0, self._num_layers * self._per_layer_dim]), + ) + + return hidden, per_layer_inputs def preprocess_weights( self, state_dict: dict[str, torch.Tensor] @@ -2120,25 +2179,25 @@ def preprocess_weights( state_dict[head_key] = state_dict[embed_key] renamed: dict[str, torch.Tensor] = {} + # Per-layer weight prefixes that should route to the embedding model + per_layer_prefixes = ( + "embed_tokens_per_layer.", + "per_layer_model_projection.", + "per_layer_projection_norm.", + ) for key, value in state_dict.items(): if key.startswith("language_model."): suffix = key[len("language_model.") :] if suffix.startswith("lm_head"): # lm_head lives directly under decoder (not decoder.model) renamed["decoder." + suffix] = value + elif any(suffix.startswith(p) for p in per_layer_prefixes): + # Per-layer embedding weights → embedding sub-model + renamed["embedding." + suffix] = value else: # All other text weights nest under decoder.model.* onnx_key = "decoder.model." + suffix - if suffix == "embed_tokens_per_layer.weight": - # HF stores one [V, L*D] weight; split into L separate - # [V, D] tables matching our nn.ModuleList layout. - num_layers = self.config.num_hidden_layers - per_layer_dim = self.decoder.model._per_layer_dim - for i in range(num_layers): - shard = value[:, i * per_layer_dim : (i + 1) * per_layer_dim] - renamed[f"decoder.model.embed_tokens_per_layer.{i}.weight"] = shard - else: - renamed[onnx_key] = value + renamed[onnx_key] = value if suffix == "embed_tokens.weight": # Token embedding is shared with the embedding sub-model renamed["embedding.embed_tokens.weight"] = value diff --git a/src/mobius/tasks/_gemma4.py b/src/mobius/tasks/_gemma4.py index ae1f9aac..55de6276 100644 --- a/src/mobius/tasks/_gemma4.py +++ b/src/mobius/tasks/_gemma4.py @@ -221,13 +221,12 @@ def _build_decoder( decoder: nn.Module, config: Gemma4Config, ) -> ir.Model: - """Build text decoder: inputs_embeds + input_ids -> logits + per-layer KV cache. + """Build text decoder: inputs_embeds [+ per_layer_inputs] -> logits + KV cache. - ``input_ids`` is included alongside ``inputs_embeds`` because models with - ``hidden_size_per_layer_input > 0`` (e.g. Gemma4 E2B) need the original token - IDs to compute per-layer token embeddings that condition each decoder layer. - When ``hidden_size_per_layer_input == 0`` the tensor is passed through but has - no effect (``_compute_per_layer_inputs`` short-circuits to ``None``). + When ``hidden_size_per_layer_input > 0`` (e.g. Gemma4 E2B), the decoder + accepts precomputed ``per_layer_inputs`` from the embedding model instead + of ``input_ids``. This moves the per-layer embedding computation to the + embedding model, simplifying the decoder graph. """ batch = ir.SymbolicDim("batch") seq_len = ir.SymbolicDim("sequence_len") @@ -251,11 +250,16 @@ def _build_decoder( dtype=ir.DataType.INT64, shape=[batch, seq_len], ) - input_ids = builder.input( - "input_ids", - dtype=ir.DataType.INT64, - shape=[batch, seq_len], - ) + + per_layer_inputs_val: ir.Value | None = None + per_layer_dim = getattr(config, "hidden_size_per_layer_input", 0) + if per_layer_dim: + total_per_layer = config.num_hidden_layers * per_layer_dim + per_layer_inputs_val = builder.input( + "per_layer_inputs", + dtype=config.dtype, + shape=[batch, seq_len, total_per_layer], + ) past_key_values = _make_gemma4_kv_cache_inputs(builder, config, batch, past_seq_len) @@ -264,7 +268,7 @@ def _build_decoder( inputs_embeds=inputs_embeds, attention_mask=attention_mask, position_ids=position_ids, - input_ids=input_ids, + per_layer_inputs=per_layer_inputs_val, past_key_values=past_key_values, ) @@ -410,11 +414,20 @@ def _build_embedding( shape=[num_audio_tokens, config.hidden_size], ) - inputs_embeds = embedding( + result = embedding( op, input_ids=input_ids, image_features=image_features, audio_features=audio_features_val, ) - builder.add_output(inputs_embeds, "inputs_embeds") + + per_layer_dim = getattr(config, "hidden_size_per_layer_input", 0) + if per_layer_dim: + # Embedding model returns (inputs_embeds, per_layer_inputs) when + # per-layer input gating is enabled. + inputs_embeds, per_layer_inputs = result + builder.add_output(inputs_embeds, "inputs_embeds") + builder.add_output(per_layer_inputs, "per_layer_inputs") + else: + builder.add_output(result, "inputs_embeds") return _make_model(graph)