diff --git a/src/mobius/_builder.py b/src/mobius/_builder.py index bb2d3ea82..50bdca7ff 100644 --- a/src/mobius/_builder.py +++ b/src/mobius/_builder.py @@ -201,14 +201,17 @@ 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: + if not flags.ort_lower_opset_for_ep and execution_provider != "openvino": 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 +222,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 execution_provider == "openvino": + 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/models/gemma4.py b/src/mobius/models/gemma4.py index a53e3ac07..49a98be4d 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: @@ -435,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, @@ -1110,6 +1144,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 +1409,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().name == "openvino" + 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 +2523,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 +2741,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 +2774,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.name == "openvino" + 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 @@ -2673,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) @@ -2743,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 @@ -3587,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) @@ -3765,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 abbd799ec..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, @@ -383,6 +414,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 +429,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 +448,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..0cf5d33ff 100644 --- a/src/mobius/tasks/_gemma4.py +++ b/src/mobius/tasks/_gemma4.py @@ -538,12 +538,20 @@ def _build_decoder( graph, builder = _make_graph(name="decoder") op = builder.op + caps = ep_capabilities() + is_openvino = caps.name == "openvino" + decoder_io_dtype = ir.DataType.FLOAT if is_openvino 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 +577,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 +639,11 @@ def _build_decoder( input_ids=input_ids_val, ) + 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 + logits = logits_f32 builder.add_output(logits, "logits") if static: _register_hybrid_cache_outputs(builder, present_key_values, config) @@ -807,9 +824,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().name == "openvino" + + 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)