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/_component_manifest.py b/src/mobius/_component_manifest.py index 920c44269..b4bad7ec8 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,93 @@ def get_hf_component_sources( return {name: tuple(paths) for name, paths in resolved.items()} +def _config_ties_word_embeddings(hf_config: object) -> bool: + """Honor text-config precedence, except for an explicit quantized tie.""" + + def field(value: object, name: str) -> object | None: + if isinstance(value, Mapping): + return value.get(name) + return getattr(value, name, None) + + configs = ( + hf_config, + field(hf_config, "text_config"), + field(hf_config, "llm_config"), + field(hf_config, "language_config"), + ) + if any( + bool(field(declaration, "tie_word_embeddings")) + for config in configs + if config is not None + for declaration in ( + 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( + 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 tied embeddings only for explicit, unique component source 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 +367,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 9b023c260..3a8571db9 100644 --- a/src/mobius/_inspect.py +++ b/src/mobius/_inspect.py @@ -16,11 +16,18 @@ from __future__ import annotations -__all__ = ["ComponentInfo", "inspect_components"] +__all__ = [ + "ComponentInfo", + "SharedWeightEndpoint", + "SharedWeightInfo", + "inspect_components", +] import dataclasses import logging +from mobius._component_manifest import SharedWeightEndpoint, SharedWeightInfo + logger = logging.getLogger(__name__) @@ -45,11 +52,17 @@ 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. Automatic inference requires an + unambiguous embedding/head pair with explicit source-path aliases; + source paths alone do not establish weight sharing. """ name: str role: str source_paths: tuple[str, ...] = () + shared_weights: tuple[SharedWeightInfo, ...] = () def _resolve_task_model_type_and_config( @@ -158,12 +171,22 @@ def inspect_components( model_type=model_type, hf_config=hf_config, ) + shared_weights = manifest.shared_weights + 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..040fc8d6c 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,112 @@ 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 == () + + +@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 + + +@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")} diff --git a/src/mobius/models/gemma4_test.py b/src/mobius/models/gemma4_test.py index abbd799ec..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: @@ -897,6 +898,67 @@ 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 = 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 + 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.""" 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