From eb70cad033ee1cc00fab837eb2a47aada3de7068 Mon Sep 17 00:00:00 2001 From: xiaoyu-work Date: Wed, 23 Sep 2026 10:49:25 -0700 Subject: [PATCH 1/3] Fix Gemma4 OpenVINO NPU export Emit an OpenVINO-compatible rank-4 attention graph, preserve opset 23 Attention semantics, and use canonical Gather extraction for per-layer inputs so NPUW can compile repeated decoder blocks. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> Copilot-Session: d7d02d41-4f2e-4bdd-ae7d-0d22fbc6d3bd --- src/mobius/_builder.py | 38 ++-- src/mobius/_builder_test.py | 11 ++ src/mobius/_execution_providers.py | 19 ++ src/mobius/models/gemma4.py | 274 ++++++++++++++++++++++++++++- src/mobius/models/gemma4_test.py | 44 +++++ src/mobius/tasks/_gemma4.py | 44 ++++- 6 files changed, 409 insertions(+), 21 deletions(-) diff --git a/src/mobius/_builder.py b/src/mobius/_builder.py index bb2d3ea82..ea85bebbf 100644 --- a/src/mobius/_builder.py +++ b/src/mobius/_builder.py @@ -201,14 +201,18 @@ def build_from_module( def _maybe_apply_opset_lowering(package: ModelPackage, execution_provider: str) -> None: """Lower default-domain opset 24 to 23 for sub-models where it is safe.""" - if not flags.ort_lower_opset_for_ep: + capabilities = ep_registry.require(execution_provider) + if not flags.ort_lower_opset_for_ep and not capabilities.requires_attention_opset23: return if execution_provider in ("default", "cpu"): return for name, model in package.items(): if "" not in model.graph.opset_imports: continue - if _graph_requires_opset24(model.graph): + functions = list(model.functions.values()) + if _graph_requires_opset24(model.graph) or any( + _graph_requires_opset24(function) for function in functions + ): logger.info( "Skipped opset→23 lowering for '%s' (EP=%s): graph uses " "opset-24-only ops (TensorScatter / Attention nonpad_kv_seqlen). " @@ -219,15 +223,27 @@ def _maybe_apply_opset_lowering(package: ModelPackage, execution_provider: str) continue original = model.graph.opset_imports[""] model.graph.opset_imports[""] = 23 - logger.warning( - "Lowered opset %d→23 for '%s' (EP=%s). " - "ORT does not yet register opset %d kernels for this EP. " - "Track https://github.com/microsoft/onnxruntime/issues/27729", - original, - name, - execution_provider, - original, - ) + for function in functions: + if "" in function.opset_imports: + function.opset_imports[""] = 23 + if capabilities.requires_attention_opset23: + logger.info( + "Lowered opset %d→23 for '%s' (EP=%s) to avoid the OpenVINO " + "opset-24 Attention mask Pad.", + original, + name, + execution_provider, + ) + else: + logger.warning( + "Lowered opset %d→23 for '%s' (EP=%s). " + "ORT does not yet register opset %d kernels for this EP. " + "Track https://github.com/microsoft/onnxruntime/issues/27729", + original, + name, + execution_provider, + original, + ) def _graph_requires_opset24(graph: ir.Graph) -> bool: diff --git a/src/mobius/_builder_test.py b/src/mobius/_builder_test.py index 354308ed1..5649d67f4 100644 --- a/src/mobius/_builder_test.py +++ b/src/mobius/_builder_test.py @@ -183,3 +183,14 @@ def test_maybe_apply_opset_lowering_skipped_when_flag_disabled( _maybe_apply_opset_lowering(pkg, execution_provider="cuda") assert pkg["embedding"].graph.opset_imports[""] == 24 + + +def test_maybe_apply_opset_lowering_required_by_openvino( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setattr(flags, "ort_lower_opset_for_ep", False) + pkg = ModelPackage({"decoder": _model_with(_standard_nodes())}) + + _maybe_apply_opset_lowering(pkg, execution_provider="openvino") + + assert pkg["decoder"].graph.opset_imports[""] == 23 diff --git a/src/mobius/_execution_providers.py b/src/mobius/_execution_providers.py index 98421ab15..1dd4ef609 100644 --- a/src/mobius/_execution_providers.py +++ b/src/mobius/_execution_providers.py @@ -147,6 +147,19 @@ class EpCapabilities: graph-capture-safe. Set ``True`` only for EPs that cannot execute ``Shape`` / ``ConstantOfShape`` under graph capture (currently WebGPU). + requires_rank4_attention: Whether Gemma4 should feed rank-4 BNSH + query/key/value tensors to Attention instead of the generic + rank-3 projection layout. OpenVINO uses this topology so the ONNX + frontend can create SDPA without inserting redundant dynamic + Reshape nodes around every projection and cache tensor. + requires_attention_opset23: Whether Attention models must be declared + as opset 23. OpenVINO uses this because its opset-24 Attention + translator always pads the mask dynamically, while the equivalent + opset-23 path feeds the complete mask directly to SDPA. + requires_float32_decoder_io: Whether multimodal decoder embeddings and + logits must use FLOAT at the component boundary. OpenVINO NPUW's + VLM path requires FLOAT ``inputs_embeds`` even when the decoder + computes internally in FLOAT16. """ name: str @@ -171,6 +184,9 @@ class EpCapabilities: max_buffer_size: int | None = None layered_per_layer_inputs: bool = False requires_graph_capture_rewrite: bool = False + requires_rank4_attention: bool = False + requires_attention_opset23: bool = False + requires_float32_decoder_io: bool = False def __post_init__(self) -> None: if not self.supports_fused_rope and self.qkv_pack_dtypes: @@ -302,6 +318,9 @@ def _register_builtins() -> None: supports_skip_layer_norm=False, provider_options={"device_type": "NPU"}, layered_per_layer_inputs=True, + requires_rank4_attention=True, + requires_attention_opset23=True, + requires_float32_decoder_io=True, ), EpCapabilities( name="cpu", diff --git a/src/mobius/models/gemma4.py b/src/mobius/models/gemma4.py index a53e3ac07..2cacaf89c 100644 --- a/src/mobius/models/gemma4.py +++ b/src/mobius/models/gemma4.py @@ -154,6 +154,15 @@ def _retain_last_bias_query_row(op: OpBuilder, bias: ir.Value | None) -> ir.Valu return op.Unsqueeze(last, op.Constant(value_ints=[2])) +def _retain_last_bias_query_row_openvino( + op: OpBuilder, bias: ir.Value | None +) -> ir.Value | None: + """Narrow an OpenVINO attention bias with a fusion-friendly Slice.""" + if bias is None: + return None + return op.Slice(bias, starts=[-1], ends=[np.iinfo(np.int64).max], axes=[2]) + + def _active_quantization( quantization: QuantizationConfig | None, ) -> QuantizationConfig | None: @@ -1110,6 +1119,246 @@ def __init__( ) self.k_norm = RMSNorm(head_dim, eps=config.rms_norm_eps) + @staticmethod + def _apply_rotary_4d( + op: OpBuilder, + states: ir.Value, + position_embeddings: tuple, + head_dim: int, + ) -> ir.Value: + """Apply Gemma4's half-split RoPE directly to [B, S, H, D].""" + cos, sin = position_embeddings[:2] + cos = op.Concat(cos, cos, axis=-1) + sin = op.Concat(sin, sin, axis=-1) + cos = op.Unsqueeze(cos, [2]) + sin = op.Unsqueeze(sin, [2]) + if head_dim % 2: + raise ValueError( + f"Gemma4 OpenVINO RoPE requires an even head dimension, got {head_dim}." + ) + first, second = op.Split(states, [head_dim // 2, head_dim // 2], axis=-1, _outputs=2) + rotated = op.Concat(op.Neg(second), first, axis=-1) + left = op.Mul(states, cos) + right = op.Mul(rotated, sin) + result = op.Add(left, right) + for value in (left, right, result): + value.type = states.type + value.shape = states.shape + return result + + @staticmethod + def _broadcast_kv_heads( + op: OpBuilder, + states: ir.Value, + attention_bias: ir.Value, + num_attention_heads: int, + num_key_value_heads: int, + head_dim: int, + ) -> ir.Value: + """Broadcast BNSH K/V heads using the MQA pattern OpenVINO fuses.""" + if num_attention_heads == num_key_value_heads: + return states + if num_attention_heads % num_key_value_heads: + raise ValueError( + f"num_attention_heads ({num_attention_heads}) must be divisible by " + f"num_key_value_heads ({num_key_value_heads})." + ) + repeats = num_attention_heads // num_key_value_heads + states_5d = op.Unsqueeze(states, [2]) + batch_shape = op.Shape(attention_bias, start=0, end=1) + sequence_shape = op.Shape(attention_bias, start=3, end=4) + target_5d = op.Concat( + batch_shape, + [num_key_value_heads], + [repeats], + sequence_shape, + [head_dim], + axis=0, + ) + states_5d = op.Expand(states_5d, target_5d) + states_5d.type = states.type + states_5d.shape = ir.Shape( + [states.shape[0], num_key_value_heads, repeats, states.shape[2], head_dim] + ) + result = op.Reshape( + states_5d, + op.Concat( + batch_shape, + [num_attention_heads], + sequence_shape, + [head_dim], + axis=0, + ), + ) + result.type = states.type + result.shape = ir.Shape( + [states.shape[0], num_attention_heads, states.shape[2], head_dim] + ) + return result + + def _forward_openvino_attention( + self, + op: OpBuilder, + hidden_states: ir.Value, + attention_bias: ir.Value, + position_embeddings: tuple | None, + shared_kv_states: dict | None, + past_key_value: tuple | None, + ) -> tuple[ir.Value, tuple[ir.Value, ir.Value]]: + """Emit rank-4 Attention so OpenVINO can form SDPA without layout churn.""" + batch_dim = ( + hidden_states.shape[0] + if hidden_states.shape is not None + else "component.decoder.batch" + ) + query_sequence_dim = ( + hidden_states.shape[1] + if hidden_states.shape is not None + else (1 if self.is_kv_shared_layer else "component.decoder.sequence_len") + ) + query_states = self.q_proj(op, hidden_states) + query_states = op.Reshape( + query_states, + [0, 0, self.num_attention_heads, self.head_dim], + ) + query_states.shape = ir.Shape( + [batch_dim, query_sequence_dim, self.num_attention_heads, self.head_dim] + ) + query_states = self.q_norm(op, query_states) + if position_embeddings is not None: + query_states = self._apply_rotary_4d( + op, query_states, position_embeddings, self.head_dim + ) + query_states = op.Transpose(query_states, perm=[0, 2, 1, 3]) + query_states.type = hidden_states.type + query_states.shape = ir.Shape( + [batch_dim, self.num_attention_heads, query_sequence_dim, self.head_dim] + ) + + if self.is_kv_shared_layer: + present_key, present_value = shared_kv_states[self.kv_shared_layer_index][:2] + else: + key_raw = self.k_proj(op, hidden_states) + key_states = op.Reshape( + key_raw, + [0, 0, self.num_key_value_heads, self.head_dim], + ) + key_states.shape = ir.Shape( + [batch_dim, query_sequence_dim, self.num_key_value_heads, self.head_dim] + ) + key_states = self.k_norm(op, key_states) + if position_embeddings is not None: + key_states = self._apply_rotary_4d( + op, key_states, position_embeddings, self.head_dim + ) + key_states = op.Transpose(key_states, perm=[0, 2, 1, 3]) + key_states.type = hidden_states.type + key_states.shape = ir.Shape( + [batch_dim, self.num_key_value_heads, query_sequence_dim, self.head_dim] + ) + + value_raw = ( + key_raw if self._use_alternative_attention else self.v_proj(op, hidden_states) + ) + value_states = op.Reshape( + value_raw, + [0, 0, self.num_key_value_heads, self.head_dim], + ) + value_states.shape = ir.Shape( + [batch_dim, query_sequence_dim, self.num_key_value_heads, self.head_dim] + ) + value_f32 = op.Cast(value_states, to=ir.DataType.FLOAT) + value_f32.shape = value_states.shape + mean_sq = op.ReduceMean(op.Mul(value_f32, value_f32), [-1], keepdims=1) + rms = op.Sqrt(op.Add(mean_sq, op.Constant(value_floats=[self._v_norm_eps]))) + value_states = op.CastLike(op.Div(value_f32, rms), value_states) + value_states.type = hidden_states.type + value_states.shape = ir.Shape( + [batch_dim, query_sequence_dim, self.num_key_value_heads, self.head_dim] + ) + value_states = op.Transpose(value_states, perm=[0, 2, 1, 3]) + value_states.type = hidden_states.type + value_states.shape = ir.Shape( + [batch_dim, self.num_key_value_heads, query_sequence_dim, self.head_dim] + ) + if past_key_value is not None: + present_key = op.Concat(past_key_value[0], key_states, axis=2) + present_value = op.Concat(past_key_value[1], value_states, axis=2) + present_sequence = "past_sequence_length + sequence_len" + else: + present_key, present_value = key_states, value_states + present_sequence = query_sequence_dim + present_shape = ir.Shape( + [batch_dim, self.num_key_value_heads, present_sequence, self.head_dim] + ) + for value in (present_key, present_value): + value.type = hidden_states.type + value.shape = present_shape + + key_states = self._broadcast_kv_heads( + op, + present_key, + attention_bias, + self.num_attention_heads, + self.num_key_value_heads, + self.head_dim, + ) + value_states = self._broadcast_kv_heads( + op, + present_value, + attention_bias, + self.num_attention_heads, + self.num_key_value_heads, + self.head_dim, + ) + attn_output = op.Attention( + query_states, + key_states, + value_states, + attention_bias, + None, + None, + scale=self.scaling, + softcap=self.softcap, + is_causal=0, + _outputs=1, + ) + attn_output.type = query_states.type + attn_output.shape = ir.Shape( + [batch_dim, self.num_attention_heads, query_sequence_dim, self.head_dim] + ) + + if ( + not self.is_kv_shared_layer + and self.provides_shared_kv + and shared_kv_states is not None + ): + shared_kv_states[self.layer_idx] = (present_key, present_value, None) + + attn_output = op.Transpose(attn_output, perm=[0, 2, 1, 3]) + attn_output.type = query_states.type + attn_output.shape = ir.Shape( + [ + query_states.shape[0], + query_states.shape[2], + self.num_attention_heads, + self.head_dim, + ] + ) + attn_output = op.Reshape( + attn_output, + [0, 0, self.num_attention_heads * self.head_dim], + ) + attn_output.type = query_states.type + attn_output.shape = ir.Shape( + [ + query_states.shape[0], + query_states.shape[2], + self.num_attention_heads * self.head_dim, + ] + ) + return self.o_proj(op, attn_output), (present_key, present_value) + def forward( self, op: OpBuilder, @@ -1135,6 +1384,21 @@ def forward( # layers receive static_cache=None and read the source layer's full buffer. is_static = static_kv_seqlen is not None + if ( + ep_capabilities().requires_rank4_attention + and not use_gqa + and not is_static + and attention_bias is not None + ): + return self._forward_openvino_attention( + op, + hidden_states, + attention_bias, + position_embeddings, + shared_kv_states, + past_key_value, + ) + # Q projection + per-head Q norm # For GQA, skip manual RoPE — the op applies it internally. query_states = self.q_proj(op, hidden_states) @@ -2234,7 +2498,7 @@ def forward( ) ) per_layer_list = [ - op.Squeeze(op.Slice(per_layer_4d, starts=[i], ends=[i + 1], axes=[2]), [2]) + op.Gather(per_layer_4d, op.Constant(value_int=i), axis=2) for i in range(num_layers) ] for layer_idx in range(self._first_kv_shared_layer, num_layers): @@ -2452,7 +2716,7 @@ def forward( # layers without truncation. if past_key_values is not None: kv_iter = iter(past_key_values) - past_kvs: list = [ + past_kvs = [ None if layer.self_attn.is_kv_shared_layer else next(kv_iter) for layer in self.layers ] @@ -2485,7 +2749,11 @@ def forward( for key, value in fallback_pos_dict.items() } fallback_bias_dict = { - key: _retain_last_bias_query_row(op, value) + key: ( + _retain_last_bias_query_row_openvino(op, value) + if caps.requires_rank4_attention + else _retain_last_bias_query_row(op, value) + ) for key, value in fallback_bias_dict.items() } per_layer_input = per_layer_list[i] if per_layer_list is not None else None diff --git a/src/mobius/models/gemma4_test.py b/src/mobius/models/gemma4_test.py index abbd799ec..19dff34f3 100644 --- a/src/mobius/models/gemma4_test.py +++ b/src/mobius/models/gemma4_test.py @@ -383,6 +383,8 @@ def test_layout_matches_execution_provider(self, execution_provider, expected_ra assert len(decoder_input.shape) == expected_rank assert len(embedding_output.shape) == expected_rank if expected_rank == 4: + assert decoder_input.dtype == ir.DataType.FLOAT + assert embedding_output.dtype == ir.DataType.FLOAT assert list(decoder_input.shape[-2:]) == [ config.num_hidden_layers, config.hidden_size_per_layer_input, @@ -396,6 +398,14 @@ def test_layout_matches_execution_provider(self, execution_provider, expected_ra and any(value is decoder_input for value in node.inputs if value is not None) for node in package["decoder"].graph ) + per_layer_gathers = [ + node + for node in package["decoder"].graph + if node.op_type == "Gather" + and node.attributes.get("axis") is not None + and node.attributes["axis"].value == 2 + ] + assert len(per_layer_gathers) == config.num_hidden_layers else: assert ( decoder_input.shape[-1] @@ -407,6 +417,40 @@ def test_layout_matches_execution_provider(self, execution_provider, expected_ra ) +class TestGemma4OpenVINOAttention: + def test_openvino_emits_rank4_attention_topology(self): + from collections import Counter + + from mobius._builder import build_from_module + from mobius.tasks._gemma4 import Gemma4Task + + config = _tiny_gemma4_config( + enable_moe_block=False, + hidden_size_per_layer_input=8, + vocab_size_per_layer_input=256, + num_kv_shared_layers=0, + ) + package = build_from_module( + Gemma4Model(config), + config, + task=Gemma4Task(), + execution_provider="openvino", + ) + counts = Counter(node.op_type for node in package["decoder"].graph) + decoder = package["decoder"].graph + inputs_embeds = next( + value for value in decoder.inputs if value.name == "inputs_embeds" + ) + logits = next(value for value in decoder.outputs if value.name == "logits") + + assert counts["Attention"] == config.num_hidden_layers + assert counts["RotaryEmbedding"] == 0 + assert counts["Softmax"] == 0 + assert counts["CumSum"] == 1 + assert inputs_embeds.dtype == ir.DataType.FLOAT + assert logits.dtype == ir.DataType.FLOAT + + class TestGemma4VisionQuantization: def test_quantize_vision_emits_matmulnbits_and_keeps_activation_clipping(self): from mobius.tasks._gemma4 import Gemma4Task diff --git a/src/mobius/tasks/_gemma4.py b/src/mobius/tasks/_gemma4.py index 0beb78c1d..11c546a8e 100644 --- a/src/mobius/tasks/_gemma4.py +++ b/src/mobius/tasks/_gemma4.py @@ -538,12 +538,21 @@ def _build_decoder( graph, builder = _make_graph(name="decoder") op = builder.op + caps = ep_capabilities() + decoder_io_dtype = ( + ir.DataType.FLOAT if caps.requires_float32_decoder_io else config.dtype + ) - inputs_embeds = builder.input( + inputs_embeds_input = builder.input( "inputs_embeds", - dtype=config.dtype, + dtype=decoder_io_dtype, shape=[batch, seq_len, config.hidden_size], ) + inputs_embeds = inputs_embeds_input + if decoder_io_dtype != config.dtype: + inputs_embeds = op.Cast(inputs_embeds_input, to=config.dtype) + inputs_embeds.type = config.dtype + inputs_embeds.shape = inputs_embeds_input.shape past_seq_len = ir.SymbolicDim("past_sequence_len") # A static-cache layer masks itself from ``write_indices`` and @@ -569,17 +578,21 @@ def _build_decoder( per_layer_inputs_val: ir.Value | None = None per_layer_dim = getattr(config, "hidden_size_per_layer_input", 0) if per_layer_dim and not config.split_per_layer_embedding: - caps = ep_capabilities() per_layer_shape = ( [batch, seq_len, config.num_hidden_layers, per_layer_dim] if caps.layered_per_layer_inputs else [batch, seq_len, config.num_hidden_layers * per_layer_dim] ) - per_layer_inputs_val = builder.input( + per_layer_inputs_input = builder.input( "per_layer_inputs", - dtype=config.dtype, + dtype=decoder_io_dtype, shape=per_layer_shape, ) + per_layer_inputs_val = per_layer_inputs_input + if decoder_io_dtype != config.dtype: + per_layer_inputs_val = op.Cast(per_layer_inputs_input, to=config.dtype) + per_layer_inputs_val.type = config.dtype + per_layer_inputs_val.shape = per_layer_inputs_input.shape # Vision-block bidirectional attention: the decoder receives the raw # ``input_ids`` (alongside ``inputs_embeds``) and derives the block @@ -627,6 +640,11 @@ def _build_decoder( input_ids=input_ids_val, ) + if caps.requires_float32_decoder_io and logits.dtype != ir.DataType.FLOAT: + logits_f32 = op.Cast(logits, to=ir.DataType.FLOAT) + logits_f32.type = ir.DataType.FLOAT + logits_f32.shape = logits.shape + logits = logits_f32 builder.add_output(logits, "logits") if static: _register_hybrid_cache_outputs(builder, present_key_values, config) @@ -807,9 +825,21 @@ def _build_embedding( # ``embedding`` returns a dict of named outputs: always # ``inputs_embeds``; optionally ``per_layer_inputs`` (per-layer gating). - builder.add_output(result["inputs_embeds"], "inputs_embeds") + cast_outputs = ep_capabilities().requires_float32_decoder_io + + def component_output(value: ir.Value) -> ir.Value: + if not cast_outputs or value.dtype == ir.DataType.FLOAT: + return value + output = op.Cast(value, to=ir.DataType.FLOAT) + output.type = ir.DataType.FLOAT + output.shape = value.shape + return output + + builder.add_output(component_output(result["inputs_embeds"]), "inputs_embeds") if "per_layer_inputs" in result: - builder.add_output(result["per_layer_inputs"], "per_layer_inputs") + builder.add_output( + component_output(result["per_layer_inputs"]), "per_layer_inputs" + ) return _make_model(graph) From 71a85e49773f0fd74e5a710400dbb7423e601d5b Mon Sep 17 00:00:00 2001 From: xiaoyu-work Date: Wed, 23 Sep 2026 11:34:59 -0700 Subject: [PATCH 2/3] Ignore redundant Gemma4 shared-KV weights Drop checkpoint K/V projection and norm tensors for layers that borrow a shared KV cache, including packed quantized sidecars that have no target module. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> Copilot-Session: d7d02d41-4f2e-4bdd-ae7d-0d22fbc6d3bd --- src/mobius/models/gemma4.py | 30 ++++++++++++++++++++++++++++++ src/mobius/models/gemma4_test.py | 31 +++++++++++++++++++++++++++++++ 2 files changed, 61 insertions(+) diff --git a/src/mobius/models/gemma4.py b/src/mobius/models/gemma4.py index 2cacaf89c..710f00109 100644 --- a/src/mobius/models/gemma4.py +++ b/src/mobius/models/gemma4.py @@ -444,6 +444,31 @@ def _dtype_safe_compress( # --------------------------------------------------------------------------- +def _drop_shared_kv_weights( + state_dict: dict[str, torch.Tensor], + config: Gemma4Config, + *, + prefix: str, +) -> None: + """Drop redundant K/V weights for layers that borrow a shared KV cache.""" + num_shared_layers = config.num_kv_shared_layers + if not num_shared_layers: + return + + first_shared_layer = config.num_hidden_layers - num_shared_layers + layers_prefix = f"{prefix}layers." + unused_suffixes = ("self_attn.k_proj.", "self_attn.v_proj.", "self_attn.k_norm.") + for key in list(state_dict): + if not key.startswith(layers_prefix): + continue + layer_and_suffix = key[len(layers_prefix) :] + layer_index, separator, suffix = layer_and_suffix.partition(".") + if not separator or not layer_index.isdigit(): + continue + if int(layer_index) >= first_shared_layer and suffix.startswith(unused_suffixes): + state_dict.pop(key) + + def _remap_moe_expert_weights( state_dict: dict[str, torch.Tensor], config: Gemma4Config, @@ -2941,6 +2966,7 @@ def preprocess_weights( # to our fused embedding table — no splitting needed. # (For WebGPU, splitting is handled by _Gemma4DecoderModel.preprocess_weights.) # Map HF expert weight names and fold router scale + _drop_shared_kv_weights(state_dict, self.config, prefix="model.") _remap_moe_expert_weights(state_dict, self.config) _split_per_layer_projection_weight(state_dict, "model.", self.config) return super().preprocess_weights(state_dict) @@ -3011,6 +3037,7 @@ def preprocess_weights( self, state_dict: dict[str, torch.Tensor] ) -> dict[str, torch.Tensor]: state_dict = vlm_decoder_weights(state_dict, tie=self.config.tie_word_embeddings) + _drop_shared_kv_weights(state_dict, self.config, prefix="model.") _split_per_layer_projection_weight(state_dict, "model.", self.config) # For WebGPU: split the fused [V, L*D] per-layer embedding into L separate [V, D] tables. per_layer_dim = self.config.hidden_size_per_layer_input @@ -3855,6 +3882,8 @@ def preprocess_weights( else: renamed[key] = value + _drop_shared_kv_weights(renamed, self.config, prefix="decoder.model.") + # Map HF expert weight names and fold router scale _remap_moe_expert_weights(renamed, self.config) @@ -4033,4 +4062,5 @@ def preprocess_weights( else: renamed[key] = value + _drop_shared_kv_weights(renamed, self.config, prefix="decoder.model.") return renamed diff --git a/src/mobius/models/gemma4_test.py b/src/mobius/models/gemma4_test.py index 19dff34f3..327683dc9 100644 --- a/src/mobius/models/gemma4_test.py +++ b/src/mobius/models/gemma4_test.py @@ -151,6 +151,37 @@ def test_quantized_moe_experts_fail_closed(self): with pytest.raises(NotImplementedError, match="Quantized Gemma4 MoE experts"): Gemma4Model(config).preprocess_weights(state_dict) + def test_shared_kv_layer_drops_redundant_kv_weights(self): + config = _tiny_gemma4_config( + enable_moe_block=False, + num_hidden_layers=2, + num_kv_shared_layers=1, + layer_types=["sliding_attention", "sliding_attention"], + ) + state_dict = { + "model.language_model.layers.1.self_attn.k_proj.weight_qweight": torch.zeros( + 64, 32, dtype=torch.uint8 + ), + "model.language_model.layers.1.self_attn.k_proj.weight_scales": torch.ones(64, 4), + "model.language_model.layers.1.self_attn.v_proj.weight_qweight": torch.zeros( + 64, 32, dtype=torch.uint8 + ), + "model.language_model.layers.1.self_attn.v_proj.weight_scales": torch.ones(64, 4), + "model.language_model.layers.1.self_attn.k_norm.weight": torch.ones(16), + "model.language_model.layers.1.self_attn.q_proj.weight_qweight": torch.zeros( + 64, 32, dtype=torch.uint8 + ), + "model.language_model.layers.1.self_attn.q_proj.weight_scales": torch.ones(64, 4), + } + + result = Gemma4Model(config).preprocess_weights(state_dict) + + assert not any( + token in key for key in result for token in ("k_proj", "v_proj", "k_norm") + ) + assert "decoder.model.layers.1.self_attn.q_proj.weight_qweight" in result + assert "decoder.model.layers.1.self_attn.q_proj.weight_scales" in result + def test_olive_quantized_decoder_sidecars_are_preprocessed(self): config = _tiny_gemma4_config( enable_moe_block=False, From cc982302b9edf260792316d860d68495f4559fc3 Mon Sep 17 00:00:00 2001 From: xiaoyu-work Date: Wed, 23 Sep 2026 11:56:18 -0700 Subject: [PATCH 3/3] Keep Gemma4 OpenVINO handling local Remove one-off requires_* capability flags and branch directly on the OpenVINO execution provider where the Gemma4 and opset workarounds are applied. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> Copilot-Session: d7d02d41-4f2e-4bdd-ae7d-0d22fbc6d3bd --- src/mobius/_builder.py | 5 ++--- src/mobius/_execution_providers.py | 19 ------------------- src/mobius/models/gemma4.py | 4 ++-- src/mobius/tasks/_gemma4.py | 9 ++++----- 4 files changed, 8 insertions(+), 29 deletions(-) diff --git a/src/mobius/_builder.py b/src/mobius/_builder.py index ea85bebbf..50bdca7ff 100644 --- a/src/mobius/_builder.py +++ b/src/mobius/_builder.py @@ -201,8 +201,7 @@ def build_from_module( def _maybe_apply_opset_lowering(package: ModelPackage, execution_provider: str) -> None: """Lower default-domain opset 24 to 23 for sub-models where it is safe.""" - capabilities = ep_registry.require(execution_provider) - if not flags.ort_lower_opset_for_ep and not capabilities.requires_attention_opset23: + if not flags.ort_lower_opset_for_ep and execution_provider != "openvino": return if execution_provider in ("default", "cpu"): return @@ -226,7 +225,7 @@ def _maybe_apply_opset_lowering(package: ModelPackage, execution_provider: str) for function in functions: if "" in function.opset_imports: function.opset_imports[""] = 23 - if capabilities.requires_attention_opset23: + if execution_provider == "openvino": logger.info( "Lowered opset %d→23 for '%s' (EP=%s) to avoid the OpenVINO " "opset-24 Attention mask Pad.", diff --git a/src/mobius/_execution_providers.py b/src/mobius/_execution_providers.py index 1dd4ef609..98421ab15 100644 --- a/src/mobius/_execution_providers.py +++ b/src/mobius/_execution_providers.py @@ -147,19 +147,6 @@ class EpCapabilities: graph-capture-safe. Set ``True`` only for EPs that cannot execute ``Shape`` / ``ConstantOfShape`` under graph capture (currently WebGPU). - requires_rank4_attention: Whether Gemma4 should feed rank-4 BNSH - query/key/value tensors to Attention instead of the generic - rank-3 projection layout. OpenVINO uses this topology so the ONNX - frontend can create SDPA without inserting redundant dynamic - Reshape nodes around every projection and cache tensor. - requires_attention_opset23: Whether Attention models must be declared - as opset 23. OpenVINO uses this because its opset-24 Attention - translator always pads the mask dynamically, while the equivalent - opset-23 path feeds the complete mask directly to SDPA. - requires_float32_decoder_io: Whether multimodal decoder embeddings and - logits must use FLOAT at the component boundary. OpenVINO NPUW's - VLM path requires FLOAT ``inputs_embeds`` even when the decoder - computes internally in FLOAT16. """ name: str @@ -184,9 +171,6 @@ class EpCapabilities: max_buffer_size: int | None = None layered_per_layer_inputs: bool = False requires_graph_capture_rewrite: bool = False - requires_rank4_attention: bool = False - requires_attention_opset23: bool = False - requires_float32_decoder_io: bool = False def __post_init__(self) -> None: if not self.supports_fused_rope and self.qkv_pack_dtypes: @@ -318,9 +302,6 @@ def _register_builtins() -> None: supports_skip_layer_norm=False, provider_options={"device_type": "NPU"}, layered_per_layer_inputs=True, - requires_rank4_attention=True, - requires_attention_opset23=True, - requires_float32_decoder_io=True, ), EpCapabilities( name="cpu", diff --git a/src/mobius/models/gemma4.py b/src/mobius/models/gemma4.py index 710f00109..49a98be4d 100644 --- a/src/mobius/models/gemma4.py +++ b/src/mobius/models/gemma4.py @@ -1410,7 +1410,7 @@ def forward( is_static = static_kv_seqlen is not None if ( - ep_capabilities().requires_rank4_attention + ep_capabilities().name == "openvino" and not use_gqa and not is_static and attention_bias is not None @@ -2776,7 +2776,7 @@ def forward( fallback_bias_dict = { key: ( _retain_last_bias_query_row_openvino(op, value) - if caps.requires_rank4_attention + if caps.name == "openvino" else _retain_last_bias_query_row(op, value) ) for key, value in fallback_bias_dict.items() diff --git a/src/mobius/tasks/_gemma4.py b/src/mobius/tasks/_gemma4.py index 11c546a8e..0cf5d33ff 100644 --- a/src/mobius/tasks/_gemma4.py +++ b/src/mobius/tasks/_gemma4.py @@ -539,9 +539,8 @@ def _build_decoder( graph, builder = _make_graph(name="decoder") op = builder.op caps = ep_capabilities() - decoder_io_dtype = ( - ir.DataType.FLOAT if caps.requires_float32_decoder_io else config.dtype - ) + is_openvino = caps.name == "openvino" + decoder_io_dtype = ir.DataType.FLOAT if is_openvino else config.dtype inputs_embeds_input = builder.input( "inputs_embeds", @@ -640,7 +639,7 @@ def _build_decoder( input_ids=input_ids_val, ) - if caps.requires_float32_decoder_io and logits.dtype != ir.DataType.FLOAT: + if is_openvino and logits.dtype != ir.DataType.FLOAT: logits_f32 = op.Cast(logits, to=ir.DataType.FLOAT) logits_f32.type = ir.DataType.FLOAT logits_f32.shape = logits.shape @@ -825,7 +824,7 @@ def _build_embedding( # ``embedding`` returns a dict of named outputs: always # ``inputs_embeds``; optionally ``per_layer_inputs`` (per-layer gating). - cast_outputs = ep_capabilities().requires_float32_decoder_io + cast_outputs = ep_capabilities().name == "openvino" def component_output(value: ir.Value) -> ir.Value: if not cast_outputs or value.dtype == ir.DataType.FLOAT: