Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
9 changes: 8 additions & 1 deletion src/mobius/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -22,6 +22,8 @@
"CausalLMConfig",
"CausalLMTask",
"ComponentInfo",
"SharedWeightEndpoint",
"SharedWeightInfo",
"ComponentExportDisposition",
"ComponentExportReport",
"DepthAnythingConfig",
Expand Down Expand Up @@ -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 (
Expand Down
177 changes: 176 additions & 1 deletion src/mobius/_component_manifest.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,6 +8,8 @@
__all__ = [
"ComponentDescriptor",
"ComponentManifest",
"SharedWeightEndpoint",
"SharedWeightInfo",
"get_hf_component_sources",
"resolve_component_manifest",
]
Expand All @@ -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:
Expand Down Expand Up @@ -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,
Expand All @@ -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:
Expand Down Expand Up @@ -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,
*,
Expand Down Expand Up @@ -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 (),
)
25 changes: 24 additions & 1 deletion src/mobius/_inspect.py
Original file line number Diff line number Diff line change
Expand Up @@ -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__)


Expand All @@ -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(
Expand Down Expand Up @@ -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()
]
Expand Down
Loading
Loading