From 8158337bb90118367260697819cfde0c93439ed7 Mon Sep 17 00:00:00 2001 From: Xiaoyu Zhang Date: Wed, 23 Sep 2026 16:34:33 -0700 Subject: [PATCH 1/5] Describe cross-component tied weights Expose canonical and alias parameter endpoints through component inspection so optimizers can coordinate independent builds. Materialize Gemma4 packed token embeddings into the split decoder LM head during ONNX export. Signed-off-by: Xiaoyu Zhang --- src/mobius/__init__.py | 9 +- src/mobius/_inspect.py | 149 ++++++++++++++++++++++++++++++- src/mobius/_inspect_test.py | 34 +++++++ src/mobius/models/gemma4.py | 75 ++++++++++++++-- src/mobius/models/gemma4_test.py | 58 ++++++++++++ 5 files changed, 316 insertions(+), 9 deletions(-) diff --git a/src/mobius/__init__.py b/src/mobius/__init__.py index bffa8a4bc..e4c802f1d 100644 --- a/src/mobius/__init__.py +++ b/src/mobius/__init__.py @@ -22,6 +22,8 @@ "CausalLMConfig", "CausalLMTask", "ComponentInfo", + "SharedWeightEndpoint", + "SharedWeightInfo", "ComponentExportDisposition", "ComponentExportReport", "DepthAnythingConfig", @@ -113,7 +115,12 @@ from mobius._constants import OPSET_VERSION from mobius._execution_providers import EpCapabilities, ep_registry, get_ep, register_ep from mobius._export_report import ComponentExportDisposition, ComponentExportReport -from mobius._inspect import ComponentInfo, inspect_components +from mobius._inspect import ( + ComponentInfo, + SharedWeightEndpoint, + SharedWeightInfo, + inspect_components, +) from mobius._model_package import ModelPackage from mobius._optimizations import optimize_model from mobius._registry import ( diff --git a/src/mobius/_inspect.py b/src/mobius/_inspect.py index 9b023c260..dc3c99c6b 100644 --- a/src/mobius/_inspect.py +++ b/src/mobius/_inspect.py @@ -16,14 +16,102 @@ from __future__ import annotations -__all__ = ["ComponentInfo", "inspect_components"] +__all__ = [ + "ComponentInfo", + "SharedWeightEndpoint", + "SharedWeightInfo", + "inspect_components", +] import dataclasses import logging +from collections.abc import Mapping +from typing import Any, cast logger = logging.getLogger(__name__) +@dataclasses.dataclass(frozen=True) +class SharedWeightEndpoint: + """One component-local consumer of a shared HuggingFace parameter.""" + + component: str + parameter: str + + def __post_init__(self) -> None: + if not self.component: + raise ValueError("shared-weight component must not be empty") + if not self.parameter: + raise ValueError("shared-weight parameter must not be empty") + + @classmethod + def from_value(cls, value: object) -> SharedWeightEndpoint: + """Normalize a mapping or duck-typed endpoint.""" + if isinstance(value, cls): + return value + if isinstance(value, Mapping): + return cls( + component=str(value["component"]), + parameter=str(value["parameter"]), + ) + duck_value = cast(Any, value) + return cls( + component=str(duck_value.component), + parameter=str(duck_value.parameter), + ) + + +@dataclasses.dataclass(frozen=True) +class SharedWeightInfo: + """A logical HuggingFace parameter consumed by multiple components.""" + + name: str + canonical: SharedWeightEndpoint + aliases: tuple[SharedWeightEndpoint, ...] + kind: str = "parameter_alias" + + def __post_init__(self) -> None: + if not self.name: + raise ValueError("shared-weight name must not be empty") + if not self.kind: + raise ValueError(f"shared weight {self.name!r} kind must not be empty") + if not self.aliases: + raise ValueError(f"shared weight {self.name!r} must declare at least one alias") + endpoints = (self.canonical, *self.aliases) + if len(set(endpoints)) != len(endpoints): + raise ValueError(f"shared weight {self.name!r} contains duplicate endpoints") + + @classmethod + def from_value(cls, value: object) -> SharedWeightInfo: + """Normalize a mapping or duck-typed shared-weight declaration.""" + if isinstance(value, cls): + return value + if isinstance(value, Mapping): + return cls( + name=str(value["name"]), + canonical=SharedWeightEndpoint.from_value(value["canonical"]), + aliases=tuple( + SharedWeightEndpoint.from_value(alias) + for alias in value.get("aliases", ()) + ), + kind=str(value.get("kind", "parameter_alias")), + ) + duck_value = cast(Any, value) + return cls( + name=str(duck_value.name), + canonical=SharedWeightEndpoint.from_value(duck_value.canonical), + aliases=tuple( + SharedWeightEndpoint.from_value(alias) for alias in duck_value.aliases + ), + kind=str(getattr(duck_value, "kind", "parameter_alias")), + ) + + @property + def endpoints(self) -> tuple[SharedWeightEndpoint, ...]: + """Canonical endpoint followed by every alias.""" + return (self.canonical, *self.aliases) + + @dataclasses.dataclass(frozen=True) class ComponentInfo: """A single component of a model. @@ -45,11 +133,16 @@ class ComponentInfo: tuple. Empty when the component is the whole model or the layout is unknown. Tools such as Olive use these to optimize a submodule in place before exporting the full model. + shared_weights: Cross-component shared parameters involving this + component. The same immutable declaration is attached to every + participating component so each independently selected build keeps + the complete relationship. """ name: str role: str source_paths: tuple[str, ...] = () + shared_weights: tuple[SharedWeightInfo, ...] = () def _resolve_task_model_type_and_config( @@ -119,6 +212,45 @@ def _get_hf_component_sources( return get_hf_component_sources(module_class, model_type, hf_config) +def _get_hf_shared_weights( + module_class: type, + hf_config: object, +) -> tuple[SharedWeightInfo, ...]: + """Read architecture-declared cross-component shared parameters.""" + resolver = getattr(module_class, "get_hf_shared_weights", None) + if resolver is None: + return () + return tuple(SharedWeightInfo.from_value(value) for value in resolver(hf_config=hf_config)) + + +def _validate_shared_weights(shared_weights, manifest) -> None: + """Validate component ownership and runtime paths for shared parameters.""" + seen_names: set[str] = set() + for shared_weight in shared_weights: + if shared_weight.name in seen_names: + raise ValueError( + f"shared weight {shared_weight.name!r} is declared more than once" + ) + seen_names.add(shared_weight.name) + for endpoint in shared_weight.endpoints: + if endpoint.component not in manifest: + raise ValueError( + f"shared weight {shared_weight.name!r} references unknown " + f"component {endpoint.component!r}" + ) + source_paths = manifest[endpoint.component].source_paths + module_path = endpoint.parameter.rpartition(".")[0] + if source_paths and not any( + module_path == source_path or module_path.startswith(f"{source_path}.") + for source_path in source_paths + ): + raise ValueError( + f"shared weight {shared_weight.name!r} parameter " + f"{endpoint.parameter!r} is outside component " + f"{endpoint.component!r} source paths {source_paths!r}" + ) + + def inspect_components( model_id: str, task=None, @@ -158,12 +290,27 @@ def inspect_components( model_type=model_type, hf_config=hf_config, ) + shared_weights = ( + _get_hf_shared_weights(module_class, hf_config) + if module_class is not None and hf_config is not None + else () + ) + _validate_shared_weights(shared_weights, manifest) + shared_by_component = { + name: tuple( + shared_weight + for shared_weight in shared_weights + if any(endpoint.component == name for endpoint in shared_weight.endpoints) + ) + for name in manifest + } components = [ ComponentInfo( name=component.name, role=component.role, source_paths=component.source_paths, + shared_weights=shared_by_component[component.name], ) for component in manifest.values() ] diff --git a/src/mobius/_inspect_test.py b/src/mobius/_inspect_test.py index 7a0850560..aacbba6ea 100644 --- a/src/mobius/_inspect_test.py +++ b/src/mobius/_inspect_test.py @@ -10,6 +10,8 @@ import mobius from mobius._inspect import ( ComponentInfo, + SharedWeightEndpoint, + SharedWeightInfo, _get_hf_component_sources, _resolve_task_model_type_and_config, inspect_components, @@ -100,6 +102,38 @@ def test_qwen3_vl_returns_hf_source_paths(monkeypatch): assert components["embedding"].source_paths == ("model.language_model.embed_tokens",) +def test_gemma4_reports_cross_component_tied_word_embeddings(monkeypatch): + _patch_autoconfig( + monkeypatch, + SimpleNamespace( + model_type="gemma4", + tie_word_embeddings=True, + text_config=SimpleNamespace(tie_word_embeddings=True), + ), + ) + + components = {component.name: component for component in inspect_components("fake/gemma4")} + shared_weight = SharedWeightInfo( + name="word_embeddings", + kind="tied_word_embeddings", + canonical=SharedWeightEndpoint( + component="embedding", + parameter="model.language_model.embed_tokens.weight", + ), + aliases=( + SharedWeightEndpoint( + component="decoder", + parameter="lm_head.weight", + ), + ), + ) + + assert components["decoder"].shared_weights == (shared_weight,) + assert components["embedding"].shared_weights == (shared_weight,) + assert components["vision_encoder"].shared_weights == () + assert components["audio_encoder"].shared_weights == () + + def test_qwen3_tts_embedders_return_shared_hf_source_paths(monkeypatch): _patch_autoconfig(monkeypatch, SimpleNamespace(model_type="qwen3_tts")) components = {c.name: c for c in inspect_components("fake/qwen3-tts")} diff --git a/src/mobius/models/gemma4.py b/src/mobius/models/gemma4.py index a53e3ac07..39e2ba8a9 100644 --- a/src/mobius/models/gemma4.py +++ b/src/mobius/models/gemma4.py @@ -36,6 +36,7 @@ from mobius._configs import ArchitectureConfig, Gemma4Config, QuantizationConfig from mobius._weight_utils import ( is_packed_quant_key, + materialize_split_tied_olive_lm_head, preprocess_quantized_weights, vlm_decoder_weights, vlm_embedding_weights, @@ -57,7 +58,11 @@ from mobius.components._activations import get_activation from mobius.components._gemma4_audio import Gemma4AudioEncoder from mobius.components._mlp import GatedMLP -from mobius.models.base import CausalLMModel, _retain_last_sequence_token +from mobius.models.base import ( + CausalLMModel, + _retain_last_sequence_token, + effective_tie_word_embeddings, +) from mobius.models.gemma3_text import Gemma3TextScaledWordEmbedding if TYPE_CHECKING: @@ -91,6 +96,58 @@ } +def _gemma4_hf_shared_weights(*, hf_config): + """Declare the tied token table when the HuggingFace config enables it.""" + text_config = getattr(hf_config, "text_config", None) + tied = getattr( + text_config, + "tie_word_embeddings", + getattr(hf_config, "tie_word_embeddings", False), + ) + if not tied: + return () + return ( + { + "name": "word_embeddings", + "kind": "tied_word_embeddings", + "canonical": { + "component": "embedding", + "parameter": "model.language_model.embed_tokens.weight", + }, + "aliases": ( + { + "component": "decoder", + "parameter": "lm_head.weight", + }, + ), + }, + ) + + +def _materialize_gemma4_split_tied_lm_head( + state_dict: dict[str, torch.Tensor], + config: Gemma4Config, +) -> None: + """Materialize the decoder head from the canonical packed token table.""" + if not effective_tie_word_embeddings(config) or config.component_quantization is None: + return + materialize_split_tied_olive_lm_head( + state_dict, + embed_key="model.language_model.embed_tokens.weight", + head_key="lm_head.weight", + embedding_quantization=config.quantization_for_source_paths( + "embedding", + ("model.language_model.embed_tokens",), + ignored_source_names=(), + ), + head_quantization=config.quantization_for_source_paths( + "decoder", + ("lm_head",), + ignored_source_names=(), + ), + ) + + def _split_per_layer_projection_weight( state_dict: dict[str, torch.Tensor], prefix: str, @@ -3406,6 +3463,7 @@ class Gemma4Model(nn.Module): # Runtime HF ``named_modules()`` sub-trees per ONNX component. HF_COMPONENT_SOURCES: ClassVar[dict[str, tuple[str, ...]]] = _GEMMA4_COMPONENT_SOURCES + get_hf_shared_weights = staticmethod(_gemma4_hf_shared_weights) HF_COMPONENT_MODULE_ALIASES: ClassVar[dict[str, dict[str, str]]] = { "decoder": { "model": "model.language_model", @@ -3476,17 +3534,17 @@ def preprocess_weights( ``input_ids``, so ``embed_tokens`` is not a decoder initializer — the token embedding lives only in the ``embedding`` sub-model. """ + _materialize_gemma4_split_tied_lm_head(state_dict, self.config) + # Strip top-level "model." prefix used by HF multimodal checkpoints. state_dict = { (key[len("model.") :] if key.startswith("model.") else key): value for key, value in state_dict.items() } - # Synthesize lm_head from embed_tokens when weights are tied. For a - # float checkpoint this copies ``embed_tokens.weight``; for a quantized - # checkpoint the tied MatMulNBits head tensors (weight/scales/ - # zero_points) are emitted directly by the loader, so nothing to do here. - if self.config.tie_word_embeddings: + # Float checkpoints still need the tied table copied to the split head. + # Packed Olive sidecars were materialized before prefix stripping above. + if effective_tie_word_embeddings(self.config): embed_key = "language_model.embed_tokens.weight" head_key = "language_model.lm_head.weight" if head_key not in state_dict and embed_key in state_dict: @@ -3667,6 +3725,7 @@ class Gemma4UnifiedModel(nn.Module): "audio_encoder": ("model.embed_audio",), "embedding": ("model.language_model.embed_tokens",), } + get_hf_shared_weights = staticmethod(_gemma4_hf_shared_weights) HF_COMPONENT_MODULE_ALIASES: ClassVar[dict[str, dict[str, str]]] = { "decoder": { "model": "model.language_model", @@ -3718,6 +3777,8 @@ def preprocess_weights( - ``embed_{vision,audio}.*.embedding_pre_projection_norm.*`` → skip (scale-free RMSNorm, no learnable weight) """ + _materialize_gemma4_split_tied_lm_head(state_dict, self.config) + # Strip top-level "model." prefix used by HF multimodal checkpoints. state_dict = { (key[len("model.") :] if key.startswith("model.") else key): value @@ -3725,7 +3786,7 @@ def preprocess_weights( } # Synthesize lm_head from embed_tokens when weights are tied. - if self.config.tie_word_embeddings: + if effective_tie_word_embeddings(self.config): embed_key = "language_model.embed_tokens.weight" head_key = "language_model.lm_head.weight" if head_key not in state_dict and embed_key in state_dict: diff --git a/src/mobius/models/gemma4_test.py b/src/mobius/models/gemma4_test.py index abbd799ec..0ac1c4d62 100644 --- a/src/mobius/models/gemma4_test.py +++ b/src/mobius/models/gemma4_test.py @@ -897,6 +897,64 @@ def test_quantized_embedding_sidecars_route_only_to_embedding_component(self): assert result["embedding.embed_tokens.qweight"].shape == (256, 64) assert result["embedding.embed_tokens.scales"].shape == (256, 4) + def test_tied_quantized_embedding_materializes_decoder_lm_head(self): + from mobius._component_quantization import ( + configure_component_quantization, + normalize_component_quantized_weights, + ) + from mobius.tasks._gemma4 import Gemma4Task + + config = self._config() + decoder_quantization = dataclasses.replace( + config.component_quantization["decoder"], + quantize_lm_head=True, + overrides={"lm_head": QuantizationOverride(bits=8)}, + ) + embedding_quantization = dataclasses.replace( + config.component_quantization["embedding"], + quantize_embeddings=True, + ) + config = dataclasses.replace( + config, + tie_word_embeddings=True, + component_quantization={ + **config.component_quantization, + "decoder": decoder_quantization, + "embedding": embedding_quantization, + }, + ) + module = Gemma4Model(config) + task = Gemma4Task() + manifest = configure_component_quantization(module, config, task) + qweight = torch.zeros(256, 64, dtype=torch.uint8) + scales = torch.ones(256, 4) + + renamed = module.preprocess_weights( + { + "model.language_model.embed_tokens.weight_qweight": qweight, + "model.language_model.embed_tokens.weight_scales": scales, + } + ) + + assert renamed["embedding.embed_tokens.weight_qweight"] is qweight + assert renamed["embedding.embed_tokens.weight_scales"] is scales + assert renamed["decoder.lm_head.weight_qweight"] is qweight + assert renamed["decoder.lm_head.weight_scales"] is scales + + result = normalize_component_quantized_weights( + renamed, + module, + config, + ("decoder", "vision_encoder", "audio_encoder", "embedding"), + manifest=manifest, + task=task, + ) + + assert "embedding.embed_tokens.qweight" in result + assert "embedding.embed_tokens.scales" in result + assert "decoder.lm_head.weight" in result + assert "decoder.lm_head.scales" in result + class TestScaleFreeRMSNormOverflow: """V norm should handle FP16 overflow from squaring large values.""" From dad35ccde62e3bb62465521dd16f0beba36dbf0b Mon Sep 17 00:00:00 2001 From: Xiaoyu Zhang Date: Wed, 23 Sep 2026 16:45:15 -0700 Subject: [PATCH 2/5] Infer tied weights from component aliases Use Hugging Face tie_word_embeddings metadata together with explicit component source aliases to derive standard embedding/head sharing. Remove the Gemma4-specific shared-weight resolver while retaining its split export materialization. Signed-off-by: Xiaoyu Zhang --- src/mobius/_inspect.py | 86 +++++++++++++++++++++++++++++++++---- src/mobius/models/gemma4.py | 30 ------------- 2 files changed, 77 insertions(+), 39 deletions(-) diff --git a/src/mobius/_inspect.py b/src/mobius/_inspect.py index dc3c99c6b..32b2c42b6 100644 --- a/src/mobius/_inspect.py +++ b/src/mobius/_inspect.py @@ -30,6 +30,21 @@ logger = logging.getLogger(__name__) +_INPUT_EMBEDDING_MODULE_NAMES = { + "codec_embedding", + "embed_tokens", + "shared", + "text_embedding", + "tok_embeddings", +} +_OUTPUT_HEAD_MODULE_NAMES = { + "codec_head", + "lm_head", + "output", + "output_projection", + "proj_out", +} + @dataclasses.dataclass(frozen=True) class SharedWeightEndpoint: @@ -212,15 +227,70 @@ def _get_hf_component_sources( return get_hf_component_sources(module_class, model_type, hf_config) -def _get_hf_shared_weights( - module_class: type, +def _config_ties_word_embeddings(hf_config: object) -> bool: + """Whether the parent or nested text config declares tied embeddings.""" + configs = ( + hf_config, + getattr(hf_config, "text_config", None), + getattr(hf_config, "llm_config", None), + getattr(hf_config, "language_config", None), + ) + return any( + bool(getattr(config, "tie_word_embeddings", False)) + for config in configs + if config is not None + ) + + +def _aliased_component_endpoint( + manifest, + *, + role: str, + local_names: set[str], +) -> SharedWeightEndpoint | None: + """Resolve one explicitly aliased source module for a component role.""" + candidates = { + SharedWeightEndpoint( + component=component.name, + parameter=f"{source_path}.weight", + ) + for component in manifest.values() + if component.role == role + for local_path, source_path in component.source_path_aliases + if local_path.rsplit(".", 1)[-1] in local_names + } + if len(candidates) != 1: + return None + return candidates.pop() + + +def _infer_hf_shared_weights( hf_config: object, + manifest, ) -> tuple[SharedWeightInfo, ...]: - """Read architecture-declared cross-component shared parameters.""" - resolver = getattr(module_class, "get_hf_shared_weights", None) - if resolver is None: + """Infer standard tied word embeddings from config plus explicit aliases.""" + if not _config_ties_word_embeddings(hf_config): return () - return tuple(SharedWeightInfo.from_value(value) for value in resolver(hf_config=hf_config)) + canonical = _aliased_component_endpoint( + manifest, + role="embedding", + local_names=_INPUT_EMBEDDING_MODULE_NAMES, + ) + output = _aliased_component_endpoint( + manifest, + role="decoder", + local_names=_OUTPUT_HEAD_MODULE_NAMES, + ) + if canonical is None or output is None or canonical.component == output.component: + return () + return ( + SharedWeightInfo( + name="word_embeddings", + kind="tied_word_embeddings", + canonical=canonical, + aliases=(output,), + ), + ) def _validate_shared_weights(shared_weights, manifest) -> None: @@ -291,9 +361,7 @@ def inspect_components( hf_config=hf_config, ) shared_weights = ( - _get_hf_shared_weights(module_class, hf_config) - if module_class is not None and hf_config is not None - else () + _infer_hf_shared_weights(hf_config, manifest) if hf_config is not None else () ) _validate_shared_weights(shared_weights, manifest) shared_by_component = { diff --git a/src/mobius/models/gemma4.py b/src/mobius/models/gemma4.py index 39e2ba8a9..a9b6fdbea 100644 --- a/src/mobius/models/gemma4.py +++ b/src/mobius/models/gemma4.py @@ -96,34 +96,6 @@ } -def _gemma4_hf_shared_weights(*, hf_config): - """Declare the tied token table when the HuggingFace config enables it.""" - text_config = getattr(hf_config, "text_config", None) - tied = getattr( - text_config, - "tie_word_embeddings", - getattr(hf_config, "tie_word_embeddings", False), - ) - if not tied: - return () - return ( - { - "name": "word_embeddings", - "kind": "tied_word_embeddings", - "canonical": { - "component": "embedding", - "parameter": "model.language_model.embed_tokens.weight", - }, - "aliases": ( - { - "component": "decoder", - "parameter": "lm_head.weight", - }, - ), - }, - ) - - def _materialize_gemma4_split_tied_lm_head( state_dict: dict[str, torch.Tensor], config: Gemma4Config, @@ -3463,7 +3435,6 @@ class Gemma4Model(nn.Module): # Runtime HF ``named_modules()`` sub-trees per ONNX component. HF_COMPONENT_SOURCES: ClassVar[dict[str, tuple[str, ...]]] = _GEMMA4_COMPONENT_SOURCES - get_hf_shared_weights = staticmethod(_gemma4_hf_shared_weights) HF_COMPONENT_MODULE_ALIASES: ClassVar[dict[str, dict[str, str]]] = { "decoder": { "model": "model.language_model", @@ -3725,7 +3696,6 @@ class Gemma4UnifiedModel(nn.Module): "audio_encoder": ("model.embed_audio",), "embedding": ("model.language_model.embed_tokens",), } - get_hf_shared_weights = staticmethod(_gemma4_hf_shared_weights) HF_COMPONENT_MODULE_ALIASES: ClassVar[dict[str, dict[str, str]]] = { "decoder": { "model": "model.language_model", From 465370737ba62d02779283b5b7c199927add8a96 Mon Sep 17 00:00:00 2001 From: Xiaoyu Zhang Date: Wed, 23 Sep 2026 16:50:31 -0700 Subject: [PATCH 3/5] Materialize shared weights in the generic adapter Carry inferred shared-weight groups on ComponentManifest and expand packed aliases before architecture-specific preprocessing. Remove Gemma4-specific materialization so standard tied component models use the common loading path. Signed-off-by: Xiaoyu Zhang --- src/mobius/_component_manifest.py | 156 ++++++++++++++++++++- src/mobius/_inspect.py | 199 +-------------------------- src/mobius/models/gemma4.py | 45 +----- src/mobius/models/gemma4_test.py | 8 +- src/mobius/weights/_adapters.py | 39 +++++- src/mobius/weights/_adapters_test.py | 82 ++++++++++- 6 files changed, 288 insertions(+), 241 deletions(-) diff --git a/src/mobius/_component_manifest.py b/src/mobius/_component_manifest.py index 920c44269..76b025c6c 100644 --- a/src/mobius/_component_manifest.py +++ b/src/mobius/_component_manifest.py @@ -8,6 +8,8 @@ __all__ = [ "ComponentDescriptor", "ComponentManifest", + "SharedWeightEndpoint", + "SharedWeightInfo", "get_hf_component_sources", "resolve_component_manifest", ] @@ -20,6 +22,61 @@ if TYPE_CHECKING: pass +_INPUT_EMBEDDING_MODULE_NAMES = { + "codec_embedding", + "embed_tokens", + "shared", + "text_embedding", + "tok_embeddings", +} +_OUTPUT_HEAD_MODULE_NAMES = { + "codec_head", + "lm_head", + "output", + "output_projection", + "proj_out", +} + + +@dataclasses.dataclass(frozen=True) +class SharedWeightEndpoint: + """One component-local consumer of a shared HuggingFace parameter.""" + + component: str + parameter: str + + def __post_init__(self) -> None: + if not self.component: + raise ValueError("shared-weight component must not be empty") + if not self.parameter: + raise ValueError("shared-weight parameter must not be empty") + + +@dataclasses.dataclass(frozen=True) +class SharedWeightInfo: + """A logical HuggingFace parameter consumed by multiple components.""" + + name: str + canonical: SharedWeightEndpoint + aliases: tuple[SharedWeightEndpoint, ...] + kind: str = "parameter_alias" + + def __post_init__(self) -> None: + if not self.name: + raise ValueError("shared-weight name must not be empty") + if not self.kind: + raise ValueError(f"shared weight {self.name!r} kind must not be empty") + if not self.aliases: + raise ValueError(f"shared weight {self.name!r} must declare at least one alias") + endpoints = (self.canonical, *self.aliases) + if len(set(endpoints)) != len(endpoints): + raise ValueError(f"shared weight {self.name!r} contains duplicate endpoints") + + @property + def endpoints(self) -> tuple[SharedWeightEndpoint, ...]: + """Canonical endpoint followed by every alias.""" + return (self.canonical, *self.aliases) + @dataclasses.dataclass(frozen=True) class ComponentDescriptor: @@ -106,6 +163,7 @@ class ComponentManifest(Mapping[str, ComponentDescriptor]): """Ordered, immutable component metadata keyed by package component name.""" components: tuple[ComponentDescriptor, ...] + shared_weights: tuple[SharedWeightInfo, ...] = () _by_name: Mapping[str, ComponentDescriptor] = dataclasses.field( init=False, repr=False, @@ -120,6 +178,30 @@ def __post_init__(self) -> None: f"component manifest declares {component.name!r} more than once" ) by_name[component.name] = component + shared_names: set[str] = set() + for shared_weight in self.shared_weights: + if shared_weight.name in shared_names: + raise ValueError( + f"shared weight {shared_weight.name!r} is declared more than once" + ) + shared_names.add(shared_weight.name) + for endpoint in shared_weight.endpoints: + if endpoint.component not in by_name: + raise ValueError( + f"shared weight {shared_weight.name!r} references unknown " + f"component {endpoint.component!r}" + ) + source_paths = by_name[endpoint.component].source_paths + module_path = endpoint.parameter.rpartition(".")[0] + if source_paths and not any( + module_path == source_path or module_path.startswith(f"{source_path}.") + for source_path in source_paths + ): + raise ValueError( + f"shared weight {shared_weight.name!r} parameter " + f"{endpoint.parameter!r} is outside component " + f"{endpoint.component!r} source paths {source_paths!r}" + ) object.__setattr__(self, "_by_name", MappingProxyType(by_name)) def __getitem__(self, name: str) -> ComponentDescriptor: @@ -154,6 +236,72 @@ def get_hf_component_sources( return {name: tuple(paths) for name, paths in resolved.items()} +def _config_ties_word_embeddings(hf_config: object) -> bool: + """Whether the parent or nested text config declares tied embeddings.""" + configs = ( + hf_config, + getattr(hf_config, "text_config", None), + getattr(hf_config, "llm_config", None), + getattr(hf_config, "language_config", None), + ) + return any( + bool(getattr(config, "tie_word_embeddings", False)) + for config in configs + if config is not None + ) + + +def _aliased_component_endpoint( + manifest: ComponentManifest, + *, + role: str, + local_names: set[str], +) -> SharedWeightEndpoint | None: + """Resolve one explicitly aliased source module for a component role.""" + candidates = { + SharedWeightEndpoint( + component=component.name, + parameter=f"{source_path}.weight", + ) + for component in manifest.values() + if component.role == role + for local_path, source_path in component.source_path_aliases + if local_path.rsplit(".", 1)[-1] in local_names + } + if len(candidates) != 1: + return None + return candidates.pop() + + +def _infer_shared_weights( + hf_config: object, + manifest: ComponentManifest, +) -> tuple[SharedWeightInfo, ...]: + """Infer standard tied word embeddings from config plus explicit aliases.""" + if not _config_ties_word_embeddings(hf_config): + return () + canonical = _aliased_component_endpoint( + manifest, + role="embedding", + local_names=_INPUT_EMBEDDING_MODULE_NAMES, + ) + output = _aliased_component_endpoint( + manifest, + role="decoder", + local_names=_OUTPUT_HEAD_MODULE_NAMES, + ) + if canonical is None or output is None or canonical.component == output.component: + return () + return ( + SharedWeightInfo( + name="word_embeddings", + kind="tied_word_embeddings", + canonical=canonical, + aliases=(output,), + ), + ) + + def resolve_component_manifest( task: object, *, @@ -198,4 +346,10 @@ def resolve_component_manifest( ) for name in ordered_names ) - return ComponentManifest(descriptors) + manifest = ComponentManifest(descriptors) + return ComponentManifest( + descriptors, + shared_weights=_infer_shared_weights(hf_config, manifest) + if hf_config is not None + else (), + ) diff --git a/src/mobius/_inspect.py b/src/mobius/_inspect.py index 32b2c42b6..56ad4750f 100644 --- a/src/mobius/_inspect.py +++ b/src/mobius/_inspect.py @@ -25,106 +25,10 @@ import dataclasses import logging -from collections.abc import Mapping -from typing import Any, cast - -logger = logging.getLogger(__name__) - -_INPUT_EMBEDDING_MODULE_NAMES = { - "codec_embedding", - "embed_tokens", - "shared", - "text_embedding", - "tok_embeddings", -} -_OUTPUT_HEAD_MODULE_NAMES = { - "codec_head", - "lm_head", - "output", - "output_projection", - "proj_out", -} - - -@dataclasses.dataclass(frozen=True) -class SharedWeightEndpoint: - """One component-local consumer of a shared HuggingFace parameter.""" - - component: str - parameter: str - - def __post_init__(self) -> None: - if not self.component: - raise ValueError("shared-weight component must not be empty") - if not self.parameter: - raise ValueError("shared-weight parameter must not be empty") - - @classmethod - def from_value(cls, value: object) -> SharedWeightEndpoint: - """Normalize a mapping or duck-typed endpoint.""" - if isinstance(value, cls): - return value - if isinstance(value, Mapping): - return cls( - component=str(value["component"]), - parameter=str(value["parameter"]), - ) - duck_value = cast(Any, value) - return cls( - component=str(duck_value.component), - parameter=str(duck_value.parameter), - ) +from mobius._component_manifest import SharedWeightEndpoint, SharedWeightInfo -@dataclasses.dataclass(frozen=True) -class SharedWeightInfo: - """A logical HuggingFace parameter consumed by multiple components.""" - - name: str - canonical: SharedWeightEndpoint - aliases: tuple[SharedWeightEndpoint, ...] - kind: str = "parameter_alias" - - def __post_init__(self) -> None: - if not self.name: - raise ValueError("shared-weight name must not be empty") - if not self.kind: - raise ValueError(f"shared weight {self.name!r} kind must not be empty") - if not self.aliases: - raise ValueError(f"shared weight {self.name!r} must declare at least one alias") - endpoints = (self.canonical, *self.aliases) - if len(set(endpoints)) != len(endpoints): - raise ValueError(f"shared weight {self.name!r} contains duplicate endpoints") - - @classmethod - def from_value(cls, value: object) -> SharedWeightInfo: - """Normalize a mapping or duck-typed shared-weight declaration.""" - if isinstance(value, cls): - return value - if isinstance(value, Mapping): - return cls( - name=str(value["name"]), - canonical=SharedWeightEndpoint.from_value(value["canonical"]), - aliases=tuple( - SharedWeightEndpoint.from_value(alias) - for alias in value.get("aliases", ()) - ), - kind=str(value.get("kind", "parameter_alias")), - ) - duck_value = cast(Any, value) - return cls( - name=str(duck_value.name), - canonical=SharedWeightEndpoint.from_value(duck_value.canonical), - aliases=tuple( - SharedWeightEndpoint.from_value(alias) for alias in duck_value.aliases - ), - kind=str(getattr(duck_value, "kind", "parameter_alias")), - ) - - @property - def endpoints(self) -> tuple[SharedWeightEndpoint, ...]: - """Canonical endpoint followed by every alias.""" - return (self.canonical, *self.aliases) +logger = logging.getLogger(__name__) @dataclasses.dataclass(frozen=True) @@ -227,100 +131,6 @@ def _get_hf_component_sources( return get_hf_component_sources(module_class, model_type, hf_config) -def _config_ties_word_embeddings(hf_config: object) -> bool: - """Whether the parent or nested text config declares tied embeddings.""" - configs = ( - hf_config, - getattr(hf_config, "text_config", None), - getattr(hf_config, "llm_config", None), - getattr(hf_config, "language_config", None), - ) - return any( - bool(getattr(config, "tie_word_embeddings", False)) - for config in configs - if config is not None - ) - - -def _aliased_component_endpoint( - manifest, - *, - role: str, - local_names: set[str], -) -> SharedWeightEndpoint | None: - """Resolve one explicitly aliased source module for a component role.""" - candidates = { - SharedWeightEndpoint( - component=component.name, - parameter=f"{source_path}.weight", - ) - for component in manifest.values() - if component.role == role - for local_path, source_path in component.source_path_aliases - if local_path.rsplit(".", 1)[-1] in local_names - } - if len(candidates) != 1: - return None - return candidates.pop() - - -def _infer_hf_shared_weights( - hf_config: object, - manifest, -) -> tuple[SharedWeightInfo, ...]: - """Infer standard tied word embeddings from config plus explicit aliases.""" - if not _config_ties_word_embeddings(hf_config): - return () - canonical = _aliased_component_endpoint( - manifest, - role="embedding", - local_names=_INPUT_EMBEDDING_MODULE_NAMES, - ) - output = _aliased_component_endpoint( - manifest, - role="decoder", - local_names=_OUTPUT_HEAD_MODULE_NAMES, - ) - if canonical is None or output is None or canonical.component == output.component: - return () - return ( - SharedWeightInfo( - name="word_embeddings", - kind="tied_word_embeddings", - canonical=canonical, - aliases=(output,), - ), - ) - - -def _validate_shared_weights(shared_weights, manifest) -> None: - """Validate component ownership and runtime paths for shared parameters.""" - seen_names: set[str] = set() - for shared_weight in shared_weights: - if shared_weight.name in seen_names: - raise ValueError( - f"shared weight {shared_weight.name!r} is declared more than once" - ) - seen_names.add(shared_weight.name) - for endpoint in shared_weight.endpoints: - if endpoint.component not in manifest: - raise ValueError( - f"shared weight {shared_weight.name!r} references unknown " - f"component {endpoint.component!r}" - ) - source_paths = manifest[endpoint.component].source_paths - module_path = endpoint.parameter.rpartition(".")[0] - if source_paths and not any( - module_path == source_path or module_path.startswith(f"{source_path}.") - for source_path in source_paths - ): - raise ValueError( - f"shared weight {shared_weight.name!r} parameter " - f"{endpoint.parameter!r} is outside component " - f"{endpoint.component!r} source paths {source_paths!r}" - ) - - def inspect_components( model_id: str, task=None, @@ -360,10 +170,7 @@ def inspect_components( model_type=model_type, hf_config=hf_config, ) - shared_weights = ( - _infer_hf_shared_weights(hf_config, manifest) if hf_config is not None else () - ) - _validate_shared_weights(shared_weights, manifest) + shared_weights = manifest.shared_weights shared_by_component = { name: tuple( shared_weight diff --git a/src/mobius/models/gemma4.py b/src/mobius/models/gemma4.py index a9b6fdbea..a53e3ac07 100644 --- a/src/mobius/models/gemma4.py +++ b/src/mobius/models/gemma4.py @@ -36,7 +36,6 @@ from mobius._configs import ArchitectureConfig, Gemma4Config, QuantizationConfig from mobius._weight_utils import ( is_packed_quant_key, - materialize_split_tied_olive_lm_head, preprocess_quantized_weights, vlm_decoder_weights, vlm_embedding_weights, @@ -58,11 +57,7 @@ from mobius.components._activations import get_activation from mobius.components._gemma4_audio import Gemma4AudioEncoder from mobius.components._mlp import GatedMLP -from mobius.models.base import ( - CausalLMModel, - _retain_last_sequence_token, - effective_tie_word_embeddings, -) +from mobius.models.base import CausalLMModel, _retain_last_sequence_token from mobius.models.gemma3_text import Gemma3TextScaledWordEmbedding if TYPE_CHECKING: @@ -96,30 +91,6 @@ } -def _materialize_gemma4_split_tied_lm_head( - state_dict: dict[str, torch.Tensor], - config: Gemma4Config, -) -> None: - """Materialize the decoder head from the canonical packed token table.""" - if not effective_tie_word_embeddings(config) or config.component_quantization is None: - return - materialize_split_tied_olive_lm_head( - state_dict, - embed_key="model.language_model.embed_tokens.weight", - head_key="lm_head.weight", - embedding_quantization=config.quantization_for_source_paths( - "embedding", - ("model.language_model.embed_tokens",), - ignored_source_names=(), - ), - head_quantization=config.quantization_for_source_paths( - "decoder", - ("lm_head",), - ignored_source_names=(), - ), - ) - - def _split_per_layer_projection_weight( state_dict: dict[str, torch.Tensor], prefix: str, @@ -3505,17 +3476,17 @@ def preprocess_weights( ``input_ids``, so ``embed_tokens`` is not a decoder initializer — the token embedding lives only in the ``embedding`` sub-model. """ - _materialize_gemma4_split_tied_lm_head(state_dict, self.config) - # Strip top-level "model." prefix used by HF multimodal checkpoints. state_dict = { (key[len("model.") :] if key.startswith("model.") else key): value for key, value in state_dict.items() } - # Float checkpoints still need the tied table copied to the split head. - # Packed Olive sidecars were materialized before prefix stripping above. - if effective_tie_word_embeddings(self.config): + # Synthesize lm_head from embed_tokens when weights are tied. For a + # float checkpoint this copies ``embed_tokens.weight``; for a quantized + # checkpoint the tied MatMulNBits head tensors (weight/scales/ + # zero_points) are emitted directly by the loader, so nothing to do here. + if self.config.tie_word_embeddings: embed_key = "language_model.embed_tokens.weight" head_key = "language_model.lm_head.weight" if head_key not in state_dict and embed_key in state_dict: @@ -3747,8 +3718,6 @@ def preprocess_weights( - ``embed_{vision,audio}.*.embedding_pre_projection_norm.*`` → skip (scale-free RMSNorm, no learnable weight) """ - _materialize_gemma4_split_tied_lm_head(state_dict, self.config) - # Strip top-level "model." prefix used by HF multimodal checkpoints. state_dict = { (key[len("model.") :] if key.startswith("model.") else key): value @@ -3756,7 +3725,7 @@ def preprocess_weights( } # Synthesize lm_head from embed_tokens when weights are tied. - if effective_tie_word_embeddings(self.config): + if self.config.tie_word_embeddings: embed_key = "language_model.embed_tokens.weight" head_key = "language_model.lm_head.weight" if head_key not in state_dict and embed_key in state_dict: diff --git a/src/mobius/models/gemma4_test.py b/src/mobius/models/gemma4_test.py index 0ac1c4d62..e82e6397f 100644 --- a/src/mobius/models/gemma4_test.py +++ b/src/mobius/models/gemma4_test.py @@ -20,6 +20,7 @@ QuantizationOverride, ) from mobius.models.gemma4 import Gemma4CausalLMModel, Gemma4EmbeddingModel, Gemma4Model +from mobius.weights import adapt_model_weights def _tiny_gemma4_config(**overrides) -> Gemma4Config: @@ -929,11 +930,14 @@ def test_tied_quantized_embedding_materializes_decoder_lm_head(self): qweight = torch.zeros(256, 64, dtype=torch.uint8) scales = torch.ones(256, 4) - renamed = module.preprocess_weights( + renamed = adapt_model_weights( + module, { "model.language_model.embed_tokens.weight_qweight": qweight, "model.language_model.embed_tokens.weight_scales": scales, - } + }, + config=config, + manifest=manifest, ) assert renamed["embedding.embed_tokens.weight_qweight"] is qweight diff --git a/src/mobius/weights/_adapters.py b/src/mobius/weights/_adapters.py index 476a634c3..c2c2f36ce 100644 --- a/src/mobius/weights/_adapters.py +++ b/src/mobius/weights/_adapters.py @@ -20,6 +20,7 @@ from mobius._component_manifest import ComponentManifest from mobius._configs import BaseModelConfig +from mobius._weight_utils import materialize_split_tied_olive_lm_head @dataclasses.dataclass(frozen=True) @@ -42,6 +43,38 @@ def adapt( """Return semantically aligned weights without format normalization.""" +def _materialize_shared_quantized_weights( + state_dict: dict[str, torch.Tensor], + context: WeightAdapterContext, +) -> None: + """Materialize non-canonical packed aliases before model-specific routing.""" + if context.config.component_quantization is None: + return + for shared_weight in context.manifest.shared_weights: + if shared_weight.kind != "tied_word_embeddings": + continue + canonical = shared_weight.canonical + canonical_module = canonical.parameter.removesuffix(".weight") + embedding_quantization = context.config.quantization_for_source_paths( + canonical.component, + (canonical_module,), + ignored_source_names=(), + ) + for alias in shared_weight.aliases: + alias_module = alias.parameter.removesuffix(".weight") + materialize_split_tied_olive_lm_head( + state_dict, + embed_key=canonical.parameter, + head_key=alias.parameter, + embedding_quantization=embedding_quantization, + head_quantization=context.config.quantization_for_source_paths( + alias.component, + (alias_module,), + ignored_source_names=(), + ), + ) + + def adapt_model_weights( module: nn.Module, state_dict: Mapping[str, torch.Tensor], @@ -51,14 +84,16 @@ def adapt_model_weights( ) -> dict[str, torch.Tensor]: """Run an explicit adapter or the legacy ``preprocess_weights`` hook.""" context = WeightAdapterContext(config=config, manifest=manifest) + state_dict = dict(state_dict) + _materialize_shared_quantized_weights(state_dict, context) adapter: ModelWeightAdapter | None = getattr(module, "weight_adapter", None) if adapter is not None: return adapter.adapt(module, state_dict, context) preprocess = getattr(module, "preprocess_weights", None) if preprocess is None: - return dict(state_dict) - result: Any = preprocess(dict(state_dict)) + return state_dict + result: Any = preprocess(state_dict) if not isinstance(result, dict): raise TypeError( f"{type(module).__name__}.preprocess_weights must return a dict, " diff --git a/src/mobius/weights/_adapters_test.py b/src/mobius/weights/_adapters_test.py index 36a277269..f4138be5f 100644 --- a/src/mobius/weights/_adapters_test.py +++ b/src/mobius/weights/_adapters_test.py @@ -8,8 +8,13 @@ import torch from onnxscript import nn -from mobius._component_manifest import ComponentDescriptor, ComponentManifest -from mobius._configs import ArchitectureConfig +from mobius._component_manifest import ( + ComponentDescriptor, + ComponentManifest, + SharedWeightEndpoint, + SharedWeightInfo, +) +from mobius._configs import ArchitectureConfig, QuantizationConfig, QuantizationOverride from mobius.weights import WeightAdapterContext, adapt_model_weights @@ -63,3 +68,76 @@ def adapt(self, module, state_dict, context: WeightAdapterContext): assert result["adapter.weight"] is tensor assert "legacy.weight" not in result + + +def test_materializes_cross_component_tied_quantized_weight_before_adapter(): + embedding_quantization = QuantizationConfig( + bits=8, + group_size=16, + quant_method="olive", + sym=True, + quantize_embeddings=True, + ) + decoder_quantization = QuantizationConfig( + bits=4, + group_size=16, + quant_method="olive", + sym=True, + quantize_lm_head=True, + overrides={"lm_head": QuantizationOverride(bits=8)}, + ) + config = ArchitectureConfig( + tie_word_embeddings=True, + quantization=decoder_quantization, + component_quantization={ + "decoder": decoder_quantization, + "embedding": embedding_quantization, + }, + ) + manifest = ComponentManifest( + ( + ComponentDescriptor( + name="decoder", + module_attribute_path="decoder", + role="decoder", + source_paths=("lm_head",), + ), + ComponentDescriptor( + name="embedding", + module_attribute_path="embedding", + role="embedding", + source_paths=("model.embed_tokens",), + ), + ), + shared_weights=( + SharedWeightInfo( + name="word_embeddings", + kind="tied_word_embeddings", + canonical=SharedWeightEndpoint( + component="embedding", + parameter="model.embed_tokens.weight", + ), + aliases=( + SharedWeightEndpoint( + component="decoder", + parameter="lm_head.weight", + ), + ), + ), + ), + ) + qweight = torch.zeros(32, 32, dtype=torch.uint8) + scales = torch.ones(32, 2) + + result = adapt_model_weights( + nn.Module(), + { + "model.embed_tokens.weight_qweight": qweight, + "model.embed_tokens.weight_scales": scales, + }, + config=config, + manifest=manifest, + ) + + assert result["lm_head.weight_qweight"] is qweight + assert result["lm_head.weight_scales"] is scales From 3c4a8dc2892ed89960e64999b43e9519fbcc8ffe Mon Sep 17 00:00:00 2001 From: Xiaoyu Zhang Date: Fri, 25 Sep 2026 12:12:46 -0700 Subject: [PATCH 4/5] Infer shared weights from quantization tie metadata Inspect both mapping and object quantization declarations, including parsed configuration, so a model-level false flag cannot hide a tied packed embedding/head relationship. Signed-off-by: Xiaoyu Zhang --- src/mobius/_component_manifest.py | 16 +++++++++++-- src/mobius/_inspect_test.py | 40 +++++++++++++++++++++++++++++++ 2 files changed, 54 insertions(+), 2 deletions(-) diff --git a/src/mobius/_component_manifest.py b/src/mobius/_component_manifest.py index 76b025c6c..093e242d0 100644 --- a/src/mobius/_component_manifest.py +++ b/src/mobius/_component_manifest.py @@ -237,7 +237,13 @@ def get_hf_component_sources( def _config_ties_word_embeddings(hf_config: object) -> bool: - """Whether the parent or nested text config declares tied embeddings.""" + """Whether model or quantization metadata declares tied embeddings.""" + + def field(value: object, name: str) -> object | None: + if isinstance(value, Mapping): + return value.get(name) + return getattr(value, name, None) + configs = ( hf_config, getattr(hf_config, "text_config", None), @@ -245,9 +251,15 @@ def _config_ties_word_embeddings(hf_config: object) -> bool: getattr(hf_config, "language_config", None), ) return any( - bool(getattr(config, "tie_word_embeddings", False)) + bool(field(declaration, "tie_word_embeddings")) for config in configs if config is not None + for declaration in ( + config, + field(config, "quantization_config"), + field(config, "quantization"), + ) + if declaration is not None ) diff --git a/src/mobius/_inspect_test.py b/src/mobius/_inspect_test.py index aacbba6ea..1db9c6b53 100644 --- a/src/mobius/_inspect_test.py +++ b/src/mobius/_inspect_test.py @@ -134,6 +134,46 @@ def test_gemma4_reports_cross_component_tied_word_embeddings(monkeypatch): assert components["audio_encoder"].shared_weights == () +@pytest.mark.parametrize( + "quantization", + [ + {"tie_word_embeddings": True}, + SimpleNamespace(tie_word_embeddings=True), + ], +) +def test_gemma4_inspection_honors_quantization_only_tie(monkeypatch, quantization): + _patch_autoconfig( + monkeypatch, + SimpleNamespace( + model_type="gemma4", + tie_word_embeddings=False, + text_config=SimpleNamespace(tie_word_embeddings=False), + quantization_config=quantization, + ), + ) + + components = {component.name: component for component in inspect_components("fake/gemma4")} + + assert components["embedding"].shared_weights + assert components["decoder"].shared_weights == components["embedding"].shared_weights + + +def test_gemma4_inspection_honors_parsed_quantization_only_tie(monkeypatch): + _patch_autoconfig( + monkeypatch, + SimpleNamespace( + model_type="gemma4", + tie_word_embeddings=False, + text_config=SimpleNamespace(tie_word_embeddings=False), + quantization=SimpleNamespace(tie_word_embeddings=True), + ), + ) + + components = {component.name: component for component in inspect_components("fake/gemma4")} + + assert components["decoder"].shared_weights + + def test_qwen3_tts_embedders_return_shared_hf_source_paths(monkeypatch): _patch_autoconfig(monkeypatch, SimpleNamespace(model_type="qwen3_tts")) components = {c.name: c for c in inspect_components("fake/qwen3-tts")} From fb6b087cdfb07236afd548873d2eacc1dc52157a Mon Sep 17 00:00:00 2001 From: Xiaoyu Zhang Date: Fri, 25 Sep 2026 16:05:32 -0700 Subject: [PATCH 5/5] Honor effective tie precedence for split components Keep explicit quantized tying authoritative, but prefer the nested text flag over the parent flag for unquantized models. State that cross-component inference requires unique explicit source aliases and cover Gemma4 and Gemma4-unified conflicts against config parsing. Signed-off-by: Xiaoyu Zhang --- src/mobius/_component_manifest.py | 25 +++++++++++++++-------- src/mobius/_inspect.py | 5 +++-- src/mobius/_inspect_test.py | 34 +++++++++++++++++++++++++++++++ 3 files changed, 54 insertions(+), 10 deletions(-) diff --git a/src/mobius/_component_manifest.py b/src/mobius/_component_manifest.py index 093e242d0..b4bad7ec8 100644 --- a/src/mobius/_component_manifest.py +++ b/src/mobius/_component_manifest.py @@ -237,7 +237,7 @@ def get_hf_component_sources( def _config_ties_word_embeddings(hf_config: object) -> bool: - """Whether model or quantization metadata declares tied embeddings.""" + """Honor text-config precedence, except for an explicit quantized tie.""" def field(value: object, name: str) -> object | None: if isinstance(value, Mapping): @@ -246,21 +246,30 @@ def field(value: object, name: str) -> object | None: configs = ( hf_config, - getattr(hf_config, "text_config", None), - getattr(hf_config, "llm_config", None), - getattr(hf_config, "language_config", None), + field(hf_config, "text_config"), + field(hf_config, "llm_config"), + field(hf_config, "language_config"), ) - return any( + if any( bool(field(declaration, "tie_word_embeddings")) for config in configs if config is not None for declaration in ( - config, field(config, "quantization_config"), field(config, "quantization"), ) if declaration is not None - ) + ): + return True + text_config = next((config for config in configs[1:] if config is not None), None) + for config in (text_config, hf_config): + if config is None: + continue + for name in ("tie_word_embeddings", "weight_tying"): + value = field(config, name) + if value is not None: + return bool(value) + return False def _aliased_component_endpoint( @@ -289,7 +298,7 @@ def _infer_shared_weights( hf_config: object, manifest: ComponentManifest, ) -> tuple[SharedWeightInfo, ...]: - """Infer standard tied word embeddings from config plus explicit aliases.""" + """Infer tied embeddings only for explicit, unique component source aliases.""" if not _config_ties_word_embeddings(hf_config): return () canonical = _aliased_component_endpoint( diff --git a/src/mobius/_inspect.py b/src/mobius/_inspect.py index 56ad4750f..3a8571db9 100644 --- a/src/mobius/_inspect.py +++ b/src/mobius/_inspect.py @@ -54,8 +54,9 @@ class ComponentInfo: place before exporting the full model. shared_weights: Cross-component shared parameters involving this component. The same immutable declaration is attached to every - participating component so each independently selected build keeps - the complete relationship. + participating component. Automatic inference requires an + unambiguous embedding/head pair with explicit source-path aliases; + source paths alone do not establish weight sharing. """ name: str diff --git a/src/mobius/_inspect_test.py b/src/mobius/_inspect_test.py index 1db9c6b53..040fc8d6c 100644 --- a/src/mobius/_inspect_test.py +++ b/src/mobius/_inspect_test.py @@ -174,6 +174,40 @@ def test_gemma4_inspection_honors_parsed_quantization_only_tie(monkeypatch): assert components["decoder"].shared_weights +@pytest.mark.parametrize("model_type", ["gemma4", "gemma4_unified"]) +@pytest.mark.parametrize( + ("parent_tie", "text_tie", "expected"), + [ + (True, False, False), + (False, True, True), + (True, None, True), + (False, None, False), + ], +) +def test_gemma4_inspection_respects_text_tie_precedence( + monkeypatch, model_type, parent_tie, text_tie, expected +): + from mobius._configs import Gemma4Config + + hf_config = SimpleNamespace( + model_type=model_type, + tie_word_embeddings=parent_tie, + text_config=SimpleNamespace(model_type="gemma4_text", tie_word_embeddings=text_tie), + ) + _patch_autoconfig(monkeypatch, hf_config) + + components = {component.name: component for component in inspect_components("fake/gemma4")} + + assert bool(components["decoder"].shared_weights) is expected + assert components["decoder"].shared_weights == components["embedding"].shared_weights + assert ( + Gemma4Config.from_transformers( + hf_config.text_config, parent_config=hf_config + ).tie_word_embeddings + is expected + ) + + def test_qwen3_tts_embedders_return_shared_hf_source_paths(monkeypatch): _patch_autoconfig(monkeypatch, SimpleNamespace(model_type="qwen3_tts")) components = {c.name: c for c in inspect_components("fake/qwen3-tts")}