From 87e46ba4cdddf3b976061d57f353c00ff1e7e28a Mon Sep 17 00:00:00 2001 From: Xiaoyu Zhang Date: Wed, 23 Sep 2026 16:34:33 -0700 Subject: [PATCH 1/9] Preserve tied weights across component builds Consume Mobius shared-weight metadata, defer non-canonical quantization to the owning component, and restore tied quantized storage during HF assembly. Validate layouts and legacy duplicate tensors before canonicalizing them. Signed-off-by: Xiaoyu Zhang --- .../configure-workflows/build-workflow.md | 6 + olive/common/mobius_utils.py | 89 ++++- olive/model/config/model_config.py | 12 + olive/passes/pytorch/quant_utils.py | 84 +++++ olive/workflows/run/hf_component_assembly.py | 317 ++++++++++++++++-- test/common/test_mobius_utils.py | 47 ++- test/model/test_composite_model.py | 42 +++ test/passes/pytorch/test_kquant.py | 115 +++++++ test/workflows/test_hf_component_assembly.py | 209 +++++++++++- 9 files changed, 891 insertions(+), 30 deletions(-) diff --git a/docs/source/how-to/configure-workflows/build-workflow.md b/docs/source/how-to/configure-workflows/build-workflow.md index d180cd108..b2ef08856 100644 --- a/docs/source/how-to/configure-workflows/build-workflow.md +++ b/docs/source/how-to/configure-workflows/build-workflow.md @@ -179,6 +179,12 @@ By default, each named build is saved under `/`. to any other location without changing where the assembled model is saved. Olive refuses to assemble into a workflow output directory that already contains files. +When Mobius reports that parameters are tied across components, keep each component in its own build. If every +requested endpoint uses the same effective bit width, group size, and symmetry, Olive defers non-canonical aliases to +the canonical component build. Assembly keeps the canonical tensor and restores the tied-weight metadata. For legacy +artifacts that already contain every packed alias, assembly additionally requires their tensors to be identical rather +than silently tying divergent data. + The named build directories contain component-only safetensors artifacts. The workflow output contains the complete checkpoint: diff --git a/olive/common/mobius_utils.py b/olive/common/mobius_utils.py index 8c4e5c13b..4fe8fadf5 100644 --- a/olive/common/mobius_utils.py +++ b/olive/common/mobius_utils.py @@ -14,7 +14,7 @@ import logging from dataclasses import dataclass, field -from typing import Optional +from typing import Any, Optional, cast logger = logging.getLogger(__name__) @@ -28,6 +28,70 @@ def _as_path_list(value: object) -> list[str]: return [str(p) for p in value if p] +@dataclass +class SharedWeightEndpoint: + """One component-local endpoint of a shared Hugging Face parameter.""" + + component: str + parameter: str + + @classmethod + def coerce(cls, data: "SharedWeightEndpoint | dict | object") -> "SharedWeightEndpoint": + if isinstance(data, cls): + return data + if isinstance(data, dict): + return cls(component=str(data["component"]), parameter=str(data["parameter"])) + duck_data = cast("Any", data) + return cls( + component=str(duck_data.component), + parameter=str(duck_data.parameter), + ) + + def to_json(self) -> dict[str, str]: + return {"component": self.component, "parameter": self.parameter} + + +@dataclass +class SharedWeightInfo: + """A logical parameter shared by components in the source Hugging Face model.""" + + name: str + canonical: SharedWeightEndpoint + aliases: list[SharedWeightEndpoint] = field(default_factory=list) + kind: str = "parameter_alias" + + @classmethod + def coerce(cls, data: "SharedWeightInfo | dict | object") -> "SharedWeightInfo": + if isinstance(data, cls): + return data + if isinstance(data, dict): + return cls( + name=str(data["name"]), + canonical=SharedWeightEndpoint.coerce(data["canonical"]), + aliases=[SharedWeightEndpoint.coerce(alias) for alias in data.get("aliases", ())], + kind=str(data.get("kind", "parameter_alias")), + ) + duck_data = cast("Any", data) + return cls( + name=str(duck_data.name), + canonical=SharedWeightEndpoint.coerce(duck_data.canonical), + aliases=[SharedWeightEndpoint.coerce(alias) for alias in duck_data.aliases], + kind=str(getattr(duck_data, "kind", "parameter_alias")), + ) + + @property + def endpoints(self) -> list[SharedWeightEndpoint]: + return [self.canonical, *self.aliases] + + def to_json(self) -> dict: + return { + "name": self.name, + "kind": self.kind, + "canonical": self.canonical.to_json(), + "aliases": [alias.to_json() for alias in self.aliases], + } + + @dataclass class ComponentInfo: """A single component returned by a component source. @@ -47,6 +111,7 @@ class ComponentInfo: name: str role: Optional[str] = None source_paths: list[str] = field(default_factory=list) + shared_weights: list[SharedWeightInfo] = field(default_factory=list) metadata: dict = field(default_factory=dict) @classmethod @@ -65,20 +130,36 @@ def coerce(cls, data: "ComponentInfo | dict | object") -> "ComponentInfo": source_paths = data.get("source_paths") if source_paths is None: source_paths = data.get("source_path") or source.get("path") - recognized = {"name", "role", "kind", "source", "source_path", "source_paths"} + recognized = { + "name", + "role", + "kind", + "source", + "source_path", + "source_paths", + "shared_weights", + } return cls( name=data["name"], role=data.get("role") or data.get("kind"), source_paths=_as_path_list(source_paths), + shared_weights=[ + SharedWeightInfo.coerce(shared_weight) for shared_weight in data.get("shared_weights", ()) + ], metadata={k: v for k, v in data.items() if k not in recognized}, ) source_paths = getattr(data, "source_paths", None) if source_paths is None: source_paths = getattr(data, "source_path", None) + duck_data = cast("Any", data) return cls( - name=data.name, - role=getattr(data, "role", None) or getattr(data, "kind", None), + name=duck_data.name, + role=getattr(duck_data, "role", None) or getattr(duck_data, "kind", None), source_paths=_as_path_list(source_paths), + shared_weights=[ + SharedWeightInfo.coerce(shared_weight) for shared_weight in getattr(duck_data, "shared_weights", ()) + ], + metadata=dict(getattr(duck_data, "metadata", {}) or {}), ) diff --git a/olive/model/config/model_config.py b/olive/model/config/model_config.py index 62edb8844..66ee027a3 100644 --- a/olive/model/config/model_config.py +++ b/olive/model/config/model_config.py @@ -217,6 +217,18 @@ def _select_hf_component(self, names: list[str]) -> "ModelConfig": attributes["component_source_paths"] = list( dict.fromkeys(path for component in selected_components for path in component.source_paths) ) + shared_weights: dict[str, dict] = {} + for component in selected_components: + for shared_weight in component.shared_weights: + serialized = shared_weight.to_json() + existing = shared_weights.get(shared_weight.name) + if existing is not None and existing != serialized: + raise ValueError(f"Components disagree on shared weight {shared_weight.name!r}.") + shared_weights[shared_weight.name] = serialized + if shared_weights: + attributes["shared_weights"] = list(shared_weights.values()) + else: + attributes.pop("shared_weights", None) new_config["model_attributes"] = attributes return ModelConfig(type=self.type, config=new_config) diff --git a/olive/passes/pytorch/quant_utils.py b/olive/passes/pytorch/quant_utils.py index 9d7111146..4d6c06be6 100644 --- a/olive/passes/pytorch/quant_utils.py +++ b/olive/passes/pytorch/quant_utils.py @@ -364,6 +364,75 @@ def _validate_component_source_paths( ) +def _defer_shared_weight_aliases( + root_model: torch.nn.Module, + component_attributes: dict, + quantization: OliveHfQuantizationConfig, + *, + lm_head_name: str | None, + embeds_name: str | None, +) -> list[dict]: + """Defer non-canonical tied weights to their canonical component build.""" + component_name = component_attributes.get("component_name") + if not component_name: + return [] + + deferred = [] + for shared_weight in component_attributes.get("shared_weights") or (): + if shared_weight.get("kind") != "tied_word_embeddings": + continue + canonical = shared_weight["canonical"] + if canonical["component"] == component_name: + continue + alias = next( + (endpoint for endpoint in shared_weight.get("aliases") or () if endpoint["component"] == component_name), + None, + ) + if alias is None: + continue + + alias_module_name = alias["parameter"].removesuffix(".weight") + if alias_module_name == lm_head_name: + enabled = quantization.lm_head + config_field = "lm_head" + elif alias_module_name == embeds_name: + enabled = quantization.embeds + config_field = "embeds" + else: + continue + if not enabled: + continue + + canonical_module_name = canonical["parameter"].removesuffix(".weight") + canonical_module = get_attr(root_model, canonical_module_name) + alias_module = get_attr(root_model, alias_module_name) + canonical_weight = getattr(canonical_module, "weight", None) + alias_weight = getattr(alias_module, "weight", None) + if canonical_weight is None or alias_weight is None or canonical_weight is not alias_weight: + raise ValueError( + f"Shared weight {shared_weight['name']!r} metadata does not match " + f"the loaded model: {canonical['parameter']!r} and " + f"{alias['parameter']!r} are not the same parameter." + ) + + qargs = quantization.get_qlinear_init_args(alias_module_name) + deferred.append( + { + "name": shared_weight["name"], + "kind": shared_weight["kind"], + "canonical": canonical, + "alias": alias, + "quantization": { + "bits": int(qargs["bits"]), + "symmetric": bool(qargs["symmetric"]), + "group_size": int(qargs["group_size"]), + }, + } + ) + setattr(quantization, config_field, False) + return deferred + + def get_qkv_quantization_groups( wrapper: ModelWrapper, module_names: set[str] | None = None, @@ -844,6 +913,18 @@ def prepare_model( if fresh_qcfg.embeds and not component_embedding_names: raise ValueError("The selected component has no torch.nn.Embedding modules to quantize.") from None + wrapper.olive_deferred_shared_weights = ( + _defer_shared_weight_aliases( + root_model, + component_attributes, + fresh_qcfg, + lm_head_name=lm_head_name, + embeds_name=embeds_name, + ) + if existing_qcfg is None + else [] + ) + fresh_skip_patterns = list(getattr(fresh_qcfg, "modules_to_not_convert", None) or []) component_embedding_name_set = set(component_embedding_names) extra_embedding_modules = component_embedding_modules.values() if component_source_paths else () @@ -1518,6 +1599,9 @@ def finalize( save_model = wrapper.olive_root_model if wrapper.olive_root_model is not None else wrapper.model save_model.quantization_method = quant_config.quant_method save_model.config.quantization_config = quant_config + deferred_shared_weights = getattr(wrapper, "olive_deferred_shared_weights", None) + if deferred_shared_weights: + save_model.config.olive_deferred_shared_weights = deferred_shared_weights # save the quantized model — state_dict hooks drop QuantTensor entries; # only plain ``_qweight`` / ``_scales`` / ``_qzeros`` buffers diff --git a/olive/workflows/run/hf_component_assembly.py b/olive/workflows/run/hf_component_assembly.py index 40bca1e7d..7a42b43b5 100644 --- a/olive/workflows/run/hf_component_assembly.py +++ b/olive/workflows/run/hf_component_assembly.py @@ -23,6 +23,7 @@ from safetensors import safe_open from safetensors.torch import save_file +from olive.common.mobius_utils import SharedWeightEndpoint, SharedWeightInfo from olive.common.quant.hf_utils import OliveHfQuantizationConfig if TYPE_CHECKING: @@ -41,6 +42,7 @@ _QUANTIZATION_METADATA_KEYS = { "component_quantization", "olive_component_quantization", + "olive_deferred_shared_weights", "quantization_config", } _PROVENANCE_CONFIG_KEYS = {"_name_or_path", "transformers_version"} @@ -56,6 +58,19 @@ "text_embedding", "tok_embeddings", } +_INPUT_EMBEDDING_MODULE_NAMES = { + "codec_embedding", + "embed_tokens", + "shared", + "text_embedding", + "tok_embeddings", +} +_OUTPUT_HEAD_MODULE_NAMES = { + "codec_head", + "lm_head", + "output_projection", + "proj_out", +} @dataclass @@ -70,6 +85,16 @@ class _BuildArtifact: config: dict[str, Any] model_output: Any workflow_output: Any + shared_weights: list[SharedWeightInfo] + + +@dataclass +class _ResolvedSharedWeight: + info: SharedWeightInfo + canonical_artifact: _BuildArtifact + alias_artifacts: list[tuple[SharedWeightEndpoint, _BuildArtifact]] + sidecar_suffixes: tuple[str, ...] + qargs: dict[str, int | bool] class _Checkpoint: @@ -121,6 +146,29 @@ def metadata(self, key: str) -> tuple[tuple[int, ...], str]: tensor_slice = self._handles[self.key_to_path[key]].get_slice(key) return tuple(tensor_slice.get_shape()), tensor_slice.get_dtype() + def tensor_equals(self, key: str, other: _Checkpoint, other_key: str) -> bool: + """Compare two tensors exactly without materializing both in full.""" + if self.metadata(key) != other.metadata(other_key): + return False + + import torch + + shape, _ = self.metadata(key) + if not shape: + return torch.equal(self.tensor(key), other.tensor(other_key)) + + trailing_elements = 1 + for dimension in shape[1:]: + trailing_elements *= dimension + rows_per_chunk = max(1, 1_000_000 // max(trailing_elements, 1)) + left = self._handles[self.key_to_path[key]].get_slice(key) + right = other._handles[other.key_to_path[other_key]].get_slice(other_key) + for start in range(0, shape[0], rows_per_chunk): + stop = min(start + rows_per_chunk, shape[0]) + if not torch.equal(left[start:stop], right[start:stop]): + return False + return True + def _matches_source_path(key: str, source_paths: list[str]) -> bool: return any(key == path or key.startswith(f"{path}.") for path in source_paths) @@ -177,6 +225,9 @@ def _collect_build_artifacts( config=json.loads((model_dir / "config.json").read_text(encoding="utf-8")), model_output=output, workflow_output=results[build_name], + shared_weights=[ + SharedWeightInfo.coerce(shared_weight) for shared_weight in attributes.get("shared_weights", ()) + ], ) ) @@ -185,6 +236,129 @@ def _collect_build_artifacts( return artifacts +def _module_name(parameter: str) -> str: + if not parameter.endswith(".weight"): + raise ValueError(f"Shared-weight parameter {parameter!r} must identify a '.weight' parameter.") + return parameter.removesuffix(".weight") + + +def _resolve_shared_weights( + artifacts: list[_BuildArtifact], +) -> list[_ResolvedSharedWeight]: + """Resolve and validate shared quantized parameters across component builds.""" + declarations: dict[str, SharedWeightInfo] = {} + for artifact in artifacts: + for shared_weight in artifact.shared_weights: + existing = declarations.get(shared_weight.name) + if existing is not None and existing != shared_weight: + raise ValueError(f"HF component builds disagree on shared weight {shared_weight.name!r}.") + declarations[shared_weight.name] = shared_weight + + artifacts_by_component = {component: artifact for artifact in artifacts for component in artifact.components} + parsed_configs = { + artifact.name: OliveHfQuantizationConfig(**artifact.config["quantization_config"]) + for artifact in artifacts + if artifact.config.get("quantization_config") + } + resolved = [] + for shared_weight in declarations.values(): + endpoints = shared_weight.endpoints + if any(endpoint.component not in artifacts_by_component for endpoint in endpoints): + continue + + endpoint_artifacts = [(endpoint, artifacts_by_component[endpoint.component]) for endpoint in endpoints] + canonical_endpoint, canonical_artifact = endpoint_artifacts[0] + if f"{canonical_endpoint.parameter}_qweight" not in canonical_artifact.checkpoint.keys: + continue + + endpoint_qargs = [] + packed_aliases = [] + for endpoint_index, (endpoint, artifact) in enumerate(endpoint_artifacts): + qweight_present = f"{endpoint.parameter}_qweight" in artifact.checkpoint.keys + quantization = parsed_configs.get(artifact.name) + deferred_request = next( + ( + request + for request in artifact.config.get("olive_deferred_shared_weights", ()) + if request.get("name") == shared_weight.name + and request.get("alias", {}).get("parameter") == endpoint.parameter + ), + None, + ) + if endpoint_index > 0 and not qweight_present and deferred_request is None: + # Quantizing only one side intentionally breaks the source tie. + break + if qweight_present and quantization is None: + raise ValueError( + f"Shared weight {shared_weight.name!r} endpoint " + f"{endpoint.parameter!r} has packed tensors without an " + "Olive quantization_config." + ) + if qweight_present: + endpoint_qargs.append(quantization.get_qlinear_init_args(_module_name(endpoint.parameter))) + if endpoint_index > 0: + packed_aliases.append((endpoint, artifact)) + else: + endpoint_qargs.append(deferred_request["quantization"]) + if len(endpoint_qargs) != len(endpoint_artifacts): + continue + if any(qargs != endpoint_qargs[0] for qargs in endpoint_qargs[1:]): + layouts = {endpoint.parameter: qargs for (endpoint, _), qargs in zip(endpoint_artifacts, endpoint_qargs)} + raise ValueError(f"Shared weight {shared_weight.name!r} has incompatible quantization layouts: {layouts}") + + canonical_suffixes = tuple( + suffix + for suffix in ("_qweight", "_scales", "_qzeros") + if f"{canonical_endpoint.parameter}{suffix}" in canonical_artifact.checkpoint.keys + ) + if canonical_suffixes[:2] != ("_qweight", "_scales"): + raise ValueError(f"Shared weight {shared_weight.name!r} canonical endpoint is missing qweight or scales.") + for endpoint, artifact in packed_aliases: + suffixes = tuple( + suffix + for suffix in ("_qweight", "_scales", "_qzeros") + if f"{endpoint.parameter}{suffix}" in artifact.checkpoint.keys + ) + if suffixes != canonical_suffixes: + raise ValueError(f"Shared weight {shared_weight.name!r} endpoints have different packed sidecars.") + for suffix in canonical_suffixes: + canonical_key = f"{canonical_endpoint.parameter}{suffix}" + alias_key = f"{endpoint.parameter}{suffix}" + if not canonical_artifact.checkpoint.tensor_equals( + canonical_key, + artifact.checkpoint, + alias_key, + ): + raise ValueError( + f"Shared weight {shared_weight.name!r} endpoint " + f"{alias_key!r} differs from canonical tensor " + f"{canonical_key!r}." + ) + + resolved.append( + _ResolvedSharedWeight( + info=shared_weight, + canonical_artifact=canonical_artifact, + alias_artifacts=endpoint_artifacts[1:], + sidecar_suffixes=canonical_suffixes, + qargs=endpoint_qargs[0], + ) + ) + return resolved + + +def _shared_weight_exclusions( + shared_weights: list[_ResolvedSharedWeight], +) -> dict[str, set[str]]: + exclusions: dict[str, set[str]] = {} + for shared_weight in shared_weights: + for endpoint, artifact in shared_weight.alias_artifacts: + keys = exclusions.setdefault(artifact.name, set()) + keys.add(endpoint.parameter) + keys.update(f"{endpoint.parameter}{suffix}" for suffix in shared_weight.sidecar_suffixes) + return exclusions + + def _normalized_model_config(value): if isinstance(value, dict): return { @@ -265,11 +439,17 @@ def unowned_metadata(artifact: _BuildArtifact) -> dict[str, tuple[tuple[int, ... ) -def _final_checkpoint_entries(artifacts: list[_BuildArtifact]) -> list[tuple[str, _BuildArtifact]]: +def _final_checkpoint_entries( + artifacts: list[_BuildArtifact], + shared_weights: list[_ResolvedSharedWeight] | None = None, +) -> list[tuple[str, _BuildArtifact]]: + exclusions = _shared_weight_exclusions(shared_weights or []) entries = [] for artifact in artifacts: entries.extend( - (key, artifact) for key in artifact.checkpoint.keys if _matches_source_path(key, artifact.source_paths) + (key, artifact) + for key in artifact.checkpoint.keys + if _matches_source_path(key, artifact.source_paths) and key not in exclusions.get(artifact.name, set()) ) component_paths = [path for artifact in artifacts for path in artifact.source_paths] @@ -286,7 +466,11 @@ def _qweight_module_name(key: str) -> str | None: return stem.removesuffix(".weight") -def _merge_quantization_config(artifacts: list[_BuildArtifact]) -> tuple[dict | None, dict[str, Any]]: +def _merge_quantization_config( + artifacts: list[_BuildArtifact], + shared_weights: list[_ResolvedSharedWeight] | None = None, +) -> tuple[dict | None, dict[str, Any]]: + shared_weights = shared_weights or [] configs = { artifact.name: artifact.config.get("quantization_config") for artifact in artifacts @@ -296,7 +480,7 @@ def _merge_quantization_config(artifacts: list[_BuildArtifact]) -> tuple[dict | return None, {} parsed = {name: OliveHfQuantizationConfig(**config) for name, config in configs.items()} - entries = _final_checkpoint_entries(artifacts) + entries = _final_checkpoint_entries(artifacts, shared_weights) observed_args: dict[str, set[tuple[int, bool, int]]] = {} for artifact in artifacts: @@ -328,6 +512,11 @@ def _merge_quantization_config(artifacts: list[_BuildArtifact]) -> tuple[dict | ) module_args[module_name] = quant_config.get_qlinear_init_args(module_name) + for shared_weight in shared_weights: + module_args[_module_name(shared_weight.info.canonical.parameter)] = shared_weight.qargs + for alias in shared_weight.info.aliases: + module_args[_module_name(alias.parameter)] = shared_weight.qargs + if not module_args: return None, {} @@ -356,22 +545,42 @@ def _merge_quantization_config(artifacts: list[_BuildArtifact]) -> tuple[dict | float_modules.add(module_name) final_skips = [f"re:^{re.escape(module_name)}$" for module_name in sorted(float_modules)] + tied_word_embeddings = [ + shared_weight for shared_weight in shared_weights if shared_weight.info.kind == "tied_word_embeddings" + ] tying_configs = [config for config in parsed.values() if config.lm_head or config.embeds] - tying_values = {config.tie_word_embeddings for config in tying_configs} - if len(tying_values) > 1: - raise ValueError("HF component builds disagree on tied word-embedding storage.") - tie_word_embeddings = tying_values.pop() if tying_values else False - if tie_word_embeddings and not all(config.lm_head and config.embeds for config in tying_configs): - raise ValueError("Tied quantized word embeddings require both embeds and lm_head in the same build.") + if tied_word_embeddings: + tie_word_embeddings = True + else: + tying_values = {config.tie_word_embeddings for config in tying_configs} + if len(tying_values) > 1: + raise ValueError("HF component builds disagree on tied word-embedding storage.") + tie_word_embeddings = tying_values.pop() if tying_values else False + if tie_word_embeddings and not all(config.lm_head and config.embeds for config in tying_configs): + raise ValueError( + "Tied quantized word embeddings require both embeds and lm_head " + "in the same build or a declared cross-component shared weight." + ) if tie_word_embeddings and not any( module_name.rsplit(".", 1)[-1] in _WORD_EMBEDDING_MODULE_NAMES for module_name in quantized_modules ): raise ValueError("Tied word-embedding metadata does not match the assembled quantized tensors.") + shared_input_embeddings = any( + _module_name(endpoint.parameter).rsplit(".", 1)[-1] in _INPUT_EMBEDDING_MODULE_NAMES + for shared_weight in tied_word_embeddings + for endpoint in shared_weight.info.endpoints + ) + shared_output_heads = any( + _module_name(endpoint.parameter).rsplit(".", 1)[-1] in _OUTPUT_HEAD_MODULE_NAMES + for shared_weight in tied_word_embeddings + for endpoint in shared_weight.info.endpoints + ) + merged = OliveHfQuantizationConfig( **default_args, - lm_head=any(config.lm_head for config in parsed.values()), - embeds=any(config.embeds for config in parsed.values()), + lm_head=shared_output_heads or any(config.lm_head for config in parsed.values()), + embeds=shared_input_embeddings or any(config.embeds for config in parsed.values()), moe=any(config.moe for config in parsed.values()), quantize_vision=any(config.quantize_vision for config in parsed.values()), modules_to_not_convert=final_skips or None, @@ -389,15 +598,50 @@ def _merge_quantization_config(artifacts: list[_BuildArtifact]) -> tuple[dict | return merged, component_configs -def _component_quantization_mapping(artifacts: list[_BuildArtifact]) -> dict[str, dict[str, Any]]: +def _apply_shared_component_quantization( + quantization: dict[str, Any], + component: str, + shared_weights: list[_ResolvedSharedWeight], +) -> dict[str, Any]: + result = deepcopy(quantization) + for shared_weight in shared_weights: + if shared_weight.info.kind != "tied_word_embeddings": + continue + endpoints = [endpoint for endpoint in shared_weight.info.endpoints if endpoint.component == component] + if not endpoints: + continue + result["tie_word_embeddings"] = True + overrides = deepcopy(result.get("overrides") or {}) + for endpoint in endpoints: + module_name = _module_name(endpoint.parameter) + leaf = module_name.rsplit(".", 1)[-1] + if leaf in _INPUT_EMBEDDING_MODULE_NAMES: + result["embeds"] = True + if leaf in _OUTPUT_HEAD_MODULE_NAMES: + result["lm_head"] = True + override = {name: value for name, value in shared_weight.qargs.items() if result.get(name) != value} + if override: + overrides[module_name] = override + result["overrides"] = overrides or None + return result + + +def _component_quantization_mapping( + artifacts: list[_BuildArtifact], + shared_weights: list[_ResolvedSharedWeight] | None = None, +) -> dict[str, dict[str, Any]]: + shared_weights = shared_weights or [] mapping = {} for artifact in artifacts: quantization = artifact.config.get("quantization_config") if not quantization: continue - quantization = deepcopy(quantization) for component in artifact.components: - mapping[component] = quantization + mapping[component] = _apply_shared_component_quantization( + quantization, + component, + shared_weights, + ) return mapping @@ -453,17 +697,20 @@ def _copy_non_weight_files(source: Path, destination: Path) -> None: def _materialize_component_artifacts( artifacts: list[_BuildArtifact], temporary: Path, + shared_weights: list[_ResolvedSharedWeight] | None = None, ) -> tuple[dict[str, str], int, dict[str, list[str]]]: weight_map = {} total_size = 0 artifact_files = {} owned_keys = set() + exclusions = _shared_weight_exclusions(shared_weights or []) + shared_weights = shared_weights or [] for artifact in artifacts: entries = [ (key, artifact.checkpoint) for key in artifact.checkpoint.keys - if _matches_source_path(key, artifact.source_paths) + if _matches_source_path(key, artifact.source_paths) and key not in exclusions.get(artifact.name, set()) ] if not entries: raise ValueError( @@ -485,12 +732,27 @@ def _materialize_component_artifacts( weight_map.update(component_map) total_size += component_size artifact_files[artifact.name] = sorted(Path(path).name for path in component_map.values()) + quantization_config = artifact.config.get("quantization_config") + if quantization_config: + component_configs = [ + _apply_shared_component_quantization( + quantization_config, + component, + shared_weights, + ) + for component in artifact.components + ] + quantization_config = component_configs[0] + if any(config != quantization_config for config in component_configs[1:]): + raise ValueError( + f"Build {artifact.name!r} has incompatible shared-weight quantization across its components." + ) manifest = { "type": "hf_component", "components": artifact.components, "source_paths": artifact.source_paths, "passes": artifact.pass_types, - "quantization_config": artifact.config.get("quantization_config"), + "quantization_config": quantization_config, "weight_files": artifact_files[artifact.name], } (artifact_dir / "component.json").write_text(json.dumps(manifest, indent=2), encoding="utf-8") @@ -670,6 +932,7 @@ def try_assemble_hf_component_builds( with _assembly_lock(output_dir): _ensure_clean_assembly_root(output_dir) _validate_build_compatibility(artifacts) + shared_weights = _resolve_shared_weights(artifacts) logger.info( "Assembling HF component builds %s into %s", [artifact.name for artifact in artifacts], @@ -678,16 +941,23 @@ def try_assemble_hf_component_builds( with tempfile.TemporaryDirectory(prefix=".olive-hf-assembly-", dir=output_dir) as temporary_dir: temporary = Path(temporary_dir) _copy_non_weight_files(artifacts[0].model_dir, temporary) - merged_quantization, build_quantization = _merge_quantization_config(artifacts) + merged_quantization, build_quantization = _merge_quantization_config( + artifacts, + shared_weights, + ) config_path = temporary / "config.json" config = json.loads(config_path.read_text(encoding="utf-8")) + config.pop("olive_deferred_shared_weights", None) if merged_quantization is None: config.pop("quantization_config", None) else: config["quantization_config"] = merged_quantization if _changes_word_embedding_storage(artifacts): _set_tie_word_embeddings(config, merged_quantization["tie_word_embeddings"]) - component_quantization = _component_quantization_mapping(artifacts) + component_quantization = _component_quantization_mapping( + artifacts, + shared_weights, + ) if component_quantization: config["component_quantization"] = component_quantization if build_quantization: @@ -697,6 +967,7 @@ def try_assemble_hf_component_builds( weight_map, total_size, artifact_files = _materialize_component_artifacts( artifacts, temporary, + shared_weights, ) index = { "metadata": {"total_size": total_size}, @@ -707,7 +978,13 @@ def try_assemble_hf_component_builds( model_config = deepcopy(artifacts[0].model_output.olive_model_config) model_config["config"]["model_path"] = str(output_dir) attributes = dict(model_config["config"].get("model_attributes") or {}) - for name in ("component_name", "component_names", "component_role", "component_source_paths"): + for name in ( + "component_name", + "component_names", + "component_role", + "component_source_paths", + "shared_weights", + ): attributes.pop(name, None) attributes["assembled_components"] = [ component for artifact in artifacts for component in artifact.components diff --git a/test/common/test_mobius_utils.py b/test/common/test_mobius_utils.py index 872832be9..8f525c592 100644 --- a/test/common/test_mobius_utils.py +++ b/test/common/test_mobius_utils.py @@ -8,7 +8,7 @@ import pytest -from olive.common.mobius_utils import ComponentInfo, inspect_components +from olive.common.mobius_utils import ComponentInfo, SharedWeightInfo, inspect_components def test_coerce_reads_contract_dict(): @@ -36,6 +36,51 @@ def test_coerce_reads_mobius_source_paths_tuple(): assert component.source_paths == ["model.layers", "model.norm", "lm_head"] +def test_coerce_reads_cross_component_shared_weights(): + component = ComponentInfo.coerce( + types.SimpleNamespace( + name="decoder", + role="decoder", + source_paths=("model.layers", "lm_head"), + shared_weights=( + types.SimpleNamespace( + name="word_embeddings", + kind="tied_word_embeddings", + canonical=types.SimpleNamespace( + component="embedding", + parameter="model.embed_tokens.weight", + ), + aliases=( + types.SimpleNamespace( + component="decoder", + parameter="lm_head.weight", + ), + ), + ), + ), + ) + ) + + assert component.shared_weights == [ + SharedWeightInfo.coerce( + { + "name": "word_embeddings", + "kind": "tied_word_embeddings", + "canonical": { + "component": "embedding", + "parameter": "model.embed_tokens.weight", + }, + "aliases": [ + { + "component": "decoder", + "parameter": "lm_head.weight", + } + ], + } + ) + ] + + def test_coerce_falls_back_to_legacy_kind_and_source_path(): # Older mobius releases expose ``kind``/``source_path`` (singular string). component = ComponentInfo.coerce( diff --git a/test/model/test_composite_model.py b/test/model/test_composite_model.py index 081f90042..af027436d 100644 --- a/test/model/test_composite_model.py +++ b/test/model/test_composite_model.py @@ -239,6 +239,48 @@ def test_model_config_select_components_hfmodel_tags_component(monkeypatch): } +def test_model_config_select_components_hfmodel_preserves_shared_weights(monkeypatch): + shared_weight = mobius_utils.SharedWeightInfo.coerce( + { + "name": "word_embeddings", + "kind": "tied_word_embeddings", + "canonical": { + "component": "embedding", + "parameter": "model.language_model.embed_tokens.weight", + }, + "aliases": [ + { + "component": "decoder", + "parameter": "lm_head.weight", + } + ], + } + ) + monkeypatch.setattr( + mobius_utils, + "inspect_components", + lambda *args, **kwargs: [ + mobius_utils.ComponentInfo( + name="decoder", + role="decoder", + source_paths=["model.language_model.layers", "lm_head"], + shared_weights=[shared_weight], + ), + mobius_utils.ComponentInfo( + name="embedding", + role="embedding", + source_paths=["model.language_model.embed_tokens"], + shared_weights=[shared_weight], + ), + ], + ) + config = ModelConfig.model_validate({"type": "HfModel", "config": {"model_path": "some/vlm"}}) + + selected = config.select_components(["decoder"]) + + assert selected.config["model_attributes"]["shared_weights"] == [shared_weight.to_json()] + + def test_model_config_select_components_hfmodel_aggregates_multiple_components(monkeypatch): monkeypatch.setattr( mobius_utils, diff --git a/test/passes/pytorch/test_kquant.py b/test/passes/pytorch/test_kquant.py index d1c84b162..cbec2f767 100644 --- a/test/passes/pytorch/test_kquant.py +++ b/test/passes/pytorch/test_kquant.py @@ -63,6 +63,25 @@ def _make_local_tiny_qwen3_moe(save_path: Path) -> HfModelHandler: return HfModelHandler(model_path=str(save_path)) +def _make_local_tiny_tied_llama(save_path: Path) -> None: + """Save a tiny tied model for independent component quantization.""" + from transformers import LlamaConfig, LlamaForCausalLM + + torch.manual_seed(0) + save_path.mkdir(parents=True, exist_ok=True) + config = LlamaConfig( # pylint: disable=unexpected-keyword-arg + vocab_size=32, + hidden_size=32, + intermediate_size=64, + num_hidden_layers=1, + num_attention_heads=2, + num_key_value_heads=2, + tie_word_embeddings=True, + ) + LlamaForCausalLM(config).save_pretrained(save_path) + _save_trivial_tokenizer(save_path, config.vocab_size) + + def _is_quant(module: torch.nn.Module) -> bool: if not isinstance(module, (torch.nn.Linear, torch.nn.Embedding)): return False @@ -259,6 +278,102 @@ def test_kquant_consumes_selective_mixed_precision_int2_int4_int8(tmp_path: Path assert_dense_int2_mixed_precision_checkpoint(loaded, output_path) +def test_kquant_defers_noncanonical_cross_component_tied_weight(tmp_path: Path): + model_path = tmp_path / "input_model" + _make_local_tiny_tied_llama(model_path) + shared_weights = [ + { + "name": "word_embeddings", + "kind": "tied_word_embeddings", + "canonical": { + "component": "embedding", + "parameter": "model.embed_tokens.weight", + }, + "aliases": [ + { + "component": "decoder", + "parameter": "lm_head.weight", + } + ], + } + ] + decoder_model = HfModelHandler( + model_path=str(model_path), + model_attributes={ + "component_name": "decoder", + "component_role": "decoder", + "component_source_paths": [ + "model.layers", + "model.norm", + "model.rotary_emb", + "lm_head", + ], + "shared_weights": shared_weights, + }, + ) + embedding_model = HfModelHandler( + model_path=str(model_path), + model_attributes={ + "component_name": "embedding", + "component_role": "embedding", + "component_source_paths": ["model.embed_tokens"], + "shared_weights": shared_weights, + }, + ) + decoder_pass = create_pass_from_dict( + KQuant, + { + "bits": 4, + "group_size": 16, + "sym": True, + "lm_head": True, + "overrides": {"lm_head": {"bits": 8}}, + }, + disable_search=True, + ) + embedding_pass = create_pass_from_dict( + KQuant, + { + "bits": 8, + "group_size": 16, + "sym": True, + "embeds": True, + }, + disable_search=True, + ) + + decoder = decoder_pass.run(decoder_model, str(tmp_path / "decoder")).load_model() + embedding = embedding_pass.run( + embedding_model, + str(tmp_path / "embedding"), + ).load_model() + + decoder_weight = decoder.lm_head._parameters["weight"].data + embedding_weight = embedding.model.embed_tokens._parameters["weight"].data + assert not isinstance(decoder_weight, QuantTensor) + assert isinstance(embedding_weight, QuantTensor) + assert decoder.config.quantization_config.lm_head is False + assert decoder.config.olive_deferred_shared_weights == [ + { + "name": "word_embeddings", + "kind": "tied_word_embeddings", + "canonical": { + "component": "embedding", + "parameter": "model.embed_tokens.weight", + }, + "alias": { + "component": "decoder", + "parameter": "lm_head.weight", + }, + "quantization": { + "bits": 8, + "symmetric": True, + "group_size": 16, + }, + } + ] + + @pytest.mark.parametrize( ("layout", "message"), [ diff --git a/test/workflows/test_hf_component_assembly.py b/test/workflows/test_hf_component_assembly.py index 70a8c924d..acddfb5d3 100644 --- a/test/workflows/test_hf_component_assembly.py +++ b/test/workflows/test_hf_component_assembly.py @@ -69,15 +69,26 @@ def _write_checkpoint(path: Path, tensors: dict, quantization_config: dict, mode save_file(tensors, path / "model.safetensors") -def _run_config(output_dir: Path, component: str, source_path: str, pass_type: str): +def _run_config( + output_dir: Path, + component: str, + source_path: str | list[str], + pass_type: str, + *, + shared_weights: list[dict] | None = None, +): + source_paths = [source_path] if isinstance(source_path, str) else source_path + model_attributes = { + "component_name": component, + "component_source_paths": source_paths, + } + if shared_weights: + model_attributes["shared_weights"] = shared_weights return SimpleNamespace( input_model=SimpleNamespace( type="hfmodel", config={ - "model_attributes": { - "component_name": component, - "component_source_paths": [source_path], - } + "model_attributes": model_attributes, }, ), engine=SimpleNamespace(output_dir=output_dir), @@ -283,6 +294,194 @@ def test_assembles_disjoint_hf_components_and_preserves_unbuilt_weights(tmp_path assert decoder_manifest["quantization_config"]["group_size"] == 32 +@pytest.mark.parametrize( + ("failure", "message"), + [ + (None, None), + ("deferred", None), + ("layout", "incompatible quantization layouts"), + ("tensor", "differs from canonical tensor"), + ], +) +def test_assembles_cross_component_tied_word_embeddings(tmp_path, failure, message): + parent = tmp_path / "assembled" + decoder_output = parent / "decoder" + embedding_output = parent / "embedding" + decoder_model = decoder_output / "model" + embedding_model = embedding_output / "model" + decoder_output.mkdir(parents=True) + embedding_output.mkdir(parents=True) + (decoder_output / "model_config.json").write_text("{}", encoding="utf-8") + (embedding_output / "model_config.json").write_text("{}", encoding="utf-8") + + shared_weights = [ + { + "name": "word_embeddings", + "kind": "tied_word_embeddings", + "canonical": { + "component": "embedding", + "parameter": "model.language_model.embed_tokens.weight", + }, + "aliases": [ + { + "component": "decoder", + "parameter": "lm_head.weight", + } + ], + } + ] + decoder_quantization = _quantization_config( + group_size=32, + symmetric=True, + quantize_vision=False, + skips=[], + overrides=None if failure == "deferred" else {"lm_head": {"bits": 8}}, + lm_head=failure != "deferred", + ) + embedding_quantization = _quantization_config( + group_size=32, + symmetric=True, + quantize_vision=False, + skips=[], + overrides={ + "model.language_model.embed_tokens": { + "bits": 8, + **({"group_size": 16} if failure == "layout" else {}), + } + }, + embeds=True, + ) + tied_qweight = torch.arange(64, dtype=torch.uint8).reshape(8, 8) + embedding_qweight = tied_qweight.clone() + if failure == "tensor": + embedding_qweight[0, 0] += 1 + tied_scales = torch.ones(8, 2) + model_config = { + "tie_word_embeddings": False, + "text_config": {"tie_word_embeddings": False}, + } + decoder_model_config = dict(model_config) + if failure == "deferred": + decoder_model_config["olive_deferred_shared_weights"] = [ + { + "name": "word_embeddings", + "kind": "tied_word_embeddings", + "canonical": shared_weights[0]["canonical"], + "alias": shared_weights[0]["aliases"][0], + "quantization": { + "bits": 8, + "symmetric": True, + "group_size": 32, + }, + } + ] + decoder_tensors = { + "model.language_model.layers.0.weight_qweight": torch.zeros(8, 4, dtype=torch.uint8), + "model.language_model.layers.0.weight_scales": torch.ones(8, 2), + "model.unoptimized.weight": torch.ones(8, 8), + } + if failure != "deferred": + decoder_tensors.update( + { + "lm_head.weight_qweight": tied_qweight, + "lm_head.weight_scales": tied_scales, + } + ) + _write_checkpoint( + decoder_model, + decoder_tensors, + decoder_quantization, + model_config=decoder_model_config, + ) + _write_checkpoint( + embedding_model, + { + "model.language_model.embed_tokens.weight_qweight": embedding_qweight, + "model.language_model.embed_tokens.weight_scales": tied_scales, + "model.unoptimized.weight": torch.ones(8, 8), + }, + embedding_quantization, + model_config=model_config, + ) + build_configs = OrderedDict( + [ + ( + "decoder", + _run_config( + decoder_output, + "decoder", + ["model.language_model.layers", "lm_head"], + "KQuant", + shared_weights=shared_weights, + ), + ), + ( + "embedding", + _run_config( + embedding_output, + "embedding", + "model.language_model.embed_tokens", + "KQuant", + shared_weights=shared_weights, + ), + ), + ] + ) + results = OrderedDict( + [ + ("decoder", _result(decoder_model)), + ("embedding", _result(embedding_model)), + ] + ) + + if message is not None: + with pytest.raises(ValueError, match=message): + try_assemble_hf_component_builds(build_configs, results, parent) + return + + assembled = try_assemble_hf_component_builds(build_configs, results, parent) + + assert assembled == parent + assert _checkpoint_keys(parent) == { + "model.language_model.layers.0.weight_qweight", + "model.language_model.layers.0.weight_scales", + "model.language_model.embed_tokens.weight_qweight", + "model.language_model.embed_tokens.weight_scales", + "model.unoptimized.weight", + } + config = json.loads((parent / "config.json").read_text(encoding="utf-8")) + quantization = config["quantization_config"] + assert quantization["tie_word_embeddings"] is True + assert quantization["embeds"] is True + assert quantization["lm_head"] is True + merged = assembly_module.OliveHfQuantizationConfig(**quantization) + assert merged.get_qlinear_init_args("model.language_model.embed_tokens") == { + "bits": 8, + "symmetric": True, + "group_size": 32, + } + assert merged.get_qlinear_init_args("lm_head") == { + "bits": 8, + "symmetric": True, + "group_size": 32, + } + assert config["tie_word_embeddings"] is True + assert config["text_config"]["tie_word_embeddings"] is True + assert config["component_quantization"]["decoder"]["tie_word_embeddings"] is True + assert config["component_quantization"]["embedding"]["tie_word_embeddings"] is True + decoder_component = assembly_module.OliveHfQuantizationConfig(**config["component_quantization"]["decoder"]) + assert decoder_component.lm_head is True + assert decoder_component.get_qlinear_init_args("lm_head") == { + "bits": 8, + "symmetric": True, + "group_size": 32, + } + decoder_manifest = json.loads((decoder_output / "component.json").read_text(encoding="utf-8")) + embedding_manifest = json.loads((embedding_output / "component.json").read_text(encoding="utf-8")) + assert decoder_manifest["quantization_config"]["tie_word_embeddings"] is True + assert embedding_manifest["quantization_config"]["tie_word_embeddings"] is True + + def test_quantization_merge_resolves_effective_overrides_and_float_skips(): decoder_config = _quantization_config( group_size=32, From 2333b9250a20b208e13272d3717a952de290c882 Mon Sep 17 00:00:00 2001 From: Xiaoyu Zhang Date: Wed, 23 Sep 2026 17:37:18 -0700 Subject: [PATCH 2/9] Defer aliases only to selected canonical builds Propagate the workflow component set into each HF build so shared aliases are deferred only when their canonical owner is present. Otherwise honor the component pass request independently and produce an untied assembled model. Signed-off-by: Xiaoyu Zhang --- .../configure-workflows/build-workflow.md | 3 +- olive/passes/pytorch/quant_utils.py | 3 + olive/workflows/run/builds.py | 12 +++- olive/workflows/run/hf_component_assembly.py | 1 + test/passes/pytorch/test_kquant.py | 59 +++++++++++++++++++ test/workflows/test_run_builds.py | 41 +++++++++++++ 6 files changed, 115 insertions(+), 4 deletions(-) diff --git a/docs/source/how-to/configure-workflows/build-workflow.md b/docs/source/how-to/configure-workflows/build-workflow.md index b2ef08856..75f1e40ae 100644 --- a/docs/source/how-to/configure-workflows/build-workflow.md +++ b/docs/source/how-to/configure-workflows/build-workflow.md @@ -183,7 +183,8 @@ When Mobius reports that parameters are tied across components, keep each compon requested endpoint uses the same effective bit width, group size, and symmetry, Olive defers non-canonical aliases to the canonical component build. Assembly keeps the canonical tensor and restores the tied-weight metadata. For legacy artifacts that already contain every packed alias, assembly additionally requires their tensors to be identical rather -than silently tying divergent data. +than silently tying divergent data. If the workflow does not build the canonical component, Olive quantizes the +requested alias independently instead of deferring it; the assembled model is then untied. The named build directories contain component-only safetensors artifacts. The workflow output contains the complete checkpoint: diff --git a/olive/passes/pytorch/quant_utils.py b/olive/passes/pytorch/quant_utils.py index 4d6c06be6..bb3c0784f 100644 --- a/olive/passes/pytorch/quant_utils.py +++ b/olive/passes/pytorch/quant_utils.py @@ -376,6 +376,7 @@ def _defer_shared_weight_aliases( component_name = component_attributes.get("component_name") if not component_name: return [] + workflow_components = set(component_attributes.get("workflow_components") or ()) deferred = [] for shared_weight in component_attributes.get("shared_weights") or (): @@ -384,6 +385,8 @@ def _defer_shared_weight_aliases( canonical = shared_weight["canonical"] if canonical["component"] == component_name: continue + if canonical["component"] not in workflow_components: + continue alias = next( (endpoint for endpoint in shared_weight.get("aliases") or () if endpoint["component"] == component_name), None, diff --git a/olive/workflows/run/builds.py b/olive/workflows/run/builds.py index 8325f8324..00e0507e8 100644 --- a/olive/workflows/run/builds.py +++ b/olive/workflows/run/builds.py @@ -154,6 +154,9 @@ def expand_builds(run_config: dict) -> OrderedDict[str, dict]: builds = _parse_builds(raw_builds, _get_workflow_output_dir(source_config)) passes = source_config.get("passes") or {} workflow_id = source_config.get("workflow_id", DEFAULT_WORKFLOW_ID) + workflow_components = list( + dict.fromkeys(component for build in builds.values() for component in (build.components or ())) + ) expanded = OrderedDict() for build_name, build in builds.items(): @@ -177,9 +180,12 @@ def expand_builds(run_config: dict) -> OrderedDict[str, dict]: input_model = child_config.get("input_model") if input_model is None: raise ValueError(f"Build {build_name!r} selects components but no input_model is configured.") - child_config["input_model"] = ( - ModelConfig.model_validate(deepcopy(input_model)).select_components(build.components).model_dump() - ) + selected_model = ModelConfig.model_validate(deepcopy(input_model)).select_components(build.components) + if selected_model.type == "hfmodel": + attributes = dict(selected_model.config.get("model_attributes") or {}) + attributes["workflow_components"] = workflow_components + selected_model.config["model_attributes"] = attributes + child_config["input_model"] = selected_model.model_dump() expanded[build_name] = child_config diff --git a/olive/workflows/run/hf_component_assembly.py b/olive/workflows/run/hf_component_assembly.py index 7a42b43b5..bb314517d 100644 --- a/olive/workflows/run/hf_component_assembly.py +++ b/olive/workflows/run/hf_component_assembly.py @@ -984,6 +984,7 @@ def try_assemble_hf_component_builds( "component_role", "component_source_paths", "shared_weights", + "workflow_components", ): attributes.pop(name, None) attributes["assembled_components"] = [ diff --git a/test/passes/pytorch/test_kquant.py b/test/passes/pytorch/test_kquant.py index cbec2f767..91f142ec6 100644 --- a/test/passes/pytorch/test_kquant.py +++ b/test/passes/pytorch/test_kquant.py @@ -309,6 +309,7 @@ def test_kquant_defers_noncanonical_cross_component_tied_weight(tmp_path: Path): "lm_head", ], "shared_weights": shared_weights, + "workflow_components": ["decoder", "embedding"], }, ) embedding_model = HfModelHandler( @@ -318,6 +319,7 @@ def test_kquant_defers_noncanonical_cross_component_tied_weight(tmp_path: Path): "component_role": "embedding", "component_source_paths": ["model.embed_tokens"], "shared_weights": shared_weights, + "workflow_components": ["decoder", "embedding"], }, ) decoder_pass = create_pass_from_dict( @@ -374,6 +376,63 @@ def test_kquant_defers_noncanonical_cross_component_tied_weight(tmp_path: Path): ] +def test_kquant_quantizes_alias_when_canonical_component_is_not_built( + tmp_path: Path, +): + model_path = tmp_path / "input_model" + _make_local_tiny_tied_llama(model_path) + decoder_model = HfModelHandler( + model_path=str(model_path), + model_attributes={ + "component_name": "decoder", + "component_role": "decoder", + "component_source_paths": [ + "model.layers", + "model.norm", + "model.rotary_emb", + "lm_head", + ], + "workflow_components": ["decoder", "vision_encoder"], + "shared_weights": [ + { + "name": "word_embeddings", + "kind": "tied_word_embeddings", + "canonical": { + "component": "embedding", + "parameter": "model.embed_tokens.weight", + }, + "aliases": [ + { + "component": "decoder", + "parameter": "lm_head.weight", + } + ], + } + ], + }, + ) + decoder_pass = create_pass_from_dict( + KQuant, + { + "bits": 4, + "group_size": 16, + "sym": True, + "lm_head": True, + "overrides": {"lm_head": {"bits": 8}}, + }, + disable_search=True, + ) + + decoder = decoder_pass.run( + decoder_model, + str(tmp_path / "decoder"), + ).load_model() + + assert isinstance(decoder.lm_head._parameters["weight"].data, QuantTensor) + assert decoder.config.quantization_config.lm_head is True + assert not hasattr(decoder.config, "olive_deferred_shared_weights") + + @pytest.mark.parametrize( ("layout", "message"), [ diff --git a/test/workflows/test_run_builds.py b/test/workflows/test_run_builds.py index f38e100a3..972cbeced 100644 --- a/test/workflows/test_run_builds.py +++ b/test/workflows/test_run_builds.py @@ -12,7 +12,9 @@ import pytest +from olive.common import mobius_utils from olive.workflows import run as olive_run +from olive.workflows.run.builds import expand_builds from test.utils import get_pytorch_model_io_config, pytorch_model_loader # pylint: disable=attribute-defined-outside-init @@ -49,6 +51,45 @@ def _patch_engine_and_acc(self): accelerator_patch = patch.object(sys.modules[olive_run.__module__], "create_accelerator", return_value=acc_mock) return run_mock, acc_mock, engine_run_patch, accelerator_patch + def test_builds_tag_all_selected_hf_components(self, monkeypatch): + monkeypatch.setattr( + mobius_utils, + "inspect_components", + lambda *args, **kwargs: [ + mobius_utils.ComponentInfo( + name="decoder", + role="decoder", + source_paths=["model.layers", "lm_head"], + ), + mobius_utils.ComponentInfo( + name="embedding", + role="embedding", + source_paths=["model.embed_tokens"], + ), + ], + ) + config = deepcopy(self.template) + config["input_model"] = { + "type": "HfModel", + "config": {"model_path": "local/model"}, + } + config["builds"] = { + "decoder": { + "components": ["decoder"], + "pipeline": ["convert"], + }, + "embedding": { + "components": ["embedding"], + "pipeline": ["convert"], + }, + } + + expanded = expand_builds(config) + + for build in expanded.values(): + attributes = build["input_model"]["config"]["model_attributes"] + assert attributes["workflow_components"] == ["decoder", "embedding"] + def test_builds_components_on_non_composite_input_raises(self): config = deepcopy(self.template) config["builds"] = { From 4df9345c9602be5354030268b3bb0adfdb10da7f Mon Sep 17 00:00:00 2001 From: "copilot-swe-agent[bot]" <198982749+Copilot@users.noreply.github.com> Date: Thu, 24 Sep 2026 16:58:03 +0000 Subject: [PATCH 3/9] Fix pylint failures in lint CI job Co-authored-by: xiaoyu-work <85524621+xiaoyu-work@users.noreply.github.com> --- olive/common/hf/wrapper.py | 1 + olive/common/mobius_utils.py | 8 ++++---- olive/workflows/run/hf_component_assembly.py | 9 ++++++--- 3 files changed, 11 insertions(+), 7 deletions(-) diff --git a/olive/common/hf/wrapper.py b/olive/common/hf/wrapper.py index 93b9fc60d..7bbb2f891 100644 --- a/olive/common/hf/wrapper.py +++ b/olive/common/hf/wrapper.py @@ -510,6 +510,7 @@ def __init__(self, config: Union[PretrainedConfig, dict]): self.olive_component_path: Optional[str] = None self.olive_component_role: Optional[str] = None self.olive_originally_tied_embeddings = False + self.olive_deferred_shared_weights: list = [] @classmethod def _resolve_model_type(cls, config: PretrainedConfig) -> Union[str, None]: diff --git a/olive/common/mobius_utils.py b/olive/common/mobius_utils.py index 4fe8fadf5..3a9a4a586 100644 --- a/olive/common/mobius_utils.py +++ b/olive/common/mobius_utils.py @@ -14,7 +14,7 @@ import logging from dataclasses import dataclass, field -from typing import Any, Optional, cast +from typing import Any, Optional logger = logging.getLogger(__name__) @@ -41,7 +41,7 @@ def coerce(cls, data: "SharedWeightEndpoint | dict | object") -> "SharedWeightEn return data if isinstance(data, dict): return cls(component=str(data["component"]), parameter=str(data["parameter"])) - duck_data = cast("Any", data) + duck_data: Any = data return cls( component=str(duck_data.component), parameter=str(duck_data.parameter), @@ -71,7 +71,7 @@ def coerce(cls, data: "SharedWeightInfo | dict | object") -> "SharedWeightInfo": aliases=[SharedWeightEndpoint.coerce(alias) for alias in data.get("aliases", ())], kind=str(data.get("kind", "parameter_alias")), ) - duck_data = cast("Any", data) + duck_data: Any = data return cls( name=str(duck_data.name), canonical=SharedWeightEndpoint.coerce(duck_data.canonical), @@ -151,7 +151,7 @@ def coerce(cls, data: "ComponentInfo | dict | object") -> "ComponentInfo": source_paths = getattr(data, "source_paths", None) if source_paths is None: source_paths = getattr(data, "source_path", None) - duck_data = cast("Any", data) + duck_data: Any = data return cls( name=duck_data.name, role=getattr(duck_data, "role", None) or getattr(duck_data, "kind", None), diff --git a/olive/workflows/run/hf_component_assembly.py b/olive/workflows/run/hf_component_assembly.py index bb314517d..c97c3fa54 100644 --- a/olive/workflows/run/hf_component_assembly.py +++ b/olive/workflows/run/hf_component_assembly.py @@ -142,8 +142,11 @@ def keys(self) -> set[str]: def tensor(self, key: str) -> torch.Tensor: return self._handles[self.key_to_path[key]].get_tensor(key) + def tensor_slice(self, key: str): + return self._handles[self.key_to_path[key]].get_slice(key) + def metadata(self, key: str) -> tuple[tuple[int, ...], str]: - tensor_slice = self._handles[self.key_to_path[key]].get_slice(key) + tensor_slice = self.tensor_slice(key) return tuple(tensor_slice.get_shape()), tensor_slice.get_dtype() def tensor_equals(self, key: str, other: _Checkpoint, other_key: str) -> bool: @@ -161,8 +164,8 @@ def tensor_equals(self, key: str, other: _Checkpoint, other_key: str) -> bool: for dimension in shape[1:]: trailing_elements *= dimension rows_per_chunk = max(1, 1_000_000 // max(trailing_elements, 1)) - left = self._handles[self.key_to_path[key]].get_slice(key) - right = other._handles[other.key_to_path[other_key]].get_slice(other_key) + left = self.tensor_slice(key) + right = other.tensor_slice(other_key) for start in range(0, shape[0], rows_per_chunk): stop = min(start + rows_per_chunk, shape[0]) if not torch.equal(left[start:stop], right[start:stop]): From d97854ec3bc342b66b6eb9913e9a1c239c8e8368 Mon Sep 17 00:00:00 2001 From: Xiaoyu Zhang Date: Thu, 24 Sep 2026 10:44:58 -0700 Subject: [PATCH 4/9] Infer scoped HF quantization targets from component ownership Let KQuant and RTN select owned LM heads, embedding tables, and vision towers when their flags are omitted, while preserving explicit opt-outs and whole-model defaults. Plan canonical shared-table requests before deferring aliases, fail if a deferred alias has no packed owner, and cover the assembled tied and untied checkpoints. Signed-off-by: Xiaoyu Zhang --- .../configure-workflows/build-workflow.md | 26 ++- olive/passes/pytorch/kquant.py | 2 +- olive/passes/pytorch/quant_utils.py | 105 ++++++++++-- olive/passes/pytorch/rtn.py | 2 +- olive/workflows/run/builds.py | 37 +++++ olive/workflows/run/hf_component_assembly.py | 11 ++ test/passes/pytorch/test_kquant.py | 112 ++++++++++--- test/passes/pytorch/test_quant_utils.py | 27 ++++ .../passes/pytorch/test_quantization_utils.py | 3 +- test/workflows/test_hf_component_assembly.py | 118 ++++++++++++-- test/workflows/test_run_builds.py | 151 ++++++++++++++++++ 11 files changed, 544 insertions(+), 50 deletions(-) diff --git a/docs/source/how-to/configure-workflows/build-workflow.md b/docs/source/how-to/configure-workflows/build-workflow.md index 75f1e40ae..ab9faee6a 100644 --- a/docs/source/how-to/configure-workflows/build-workflow.md +++ b/docs/source/how-to/configure-workflows/build-workflow.md @@ -150,13 +150,18 @@ already contain files. "decoder_kquant": { "type": "KQuant", "bits": 4, + "group_size": 32, + "overrides": {"lm_head": {"bits": 8}} + }, + "embedding_kquant": { + "type": "KQuant", + "bits": 8, "group_size": 32 }, "vision_rtn": { "type": "Rtn", "bits": 4, - "group_size": 128, - "quantize_vision": true + "group_size": 128 } }, "engine": { @@ -167,6 +172,10 @@ already contain files. "components": ["decoder"], "pipeline": ["decoder_kquant"] }, + "embedding": { + "components": ["embedding"], + "pipeline": ["embedding_kquant"] + }, "vision": { "components": ["vision_encoder"], "pipeline": ["vision_rtn"] @@ -179,12 +188,21 @@ By default, each named build is saved under `/`. to any other location without changing where the assembled model is saved. Olive refuses to assemble into a workflow output directory that already contains files. +For KQuant and RTN on one selected Hugging Face component, omitted `lm_head`, `embeds`, and `quantize_vision` +automatically include the selected decoder's LM head, an embedding component's token tables, and owned vision-tower +weights, respectively. Explicit `false` keeps a category floating point; `modules_to_not_convert` excludes +individual modules. Whole-model passes retain their original opt-in defaults. The saved component and assembled +`quantization_config` still record `lm_head` and `embeds` as the actual checkpoint layout, including deferred +cross-component aliases. + When Mobius reports that parameters are tied across components, keep each component in its own build. If every requested endpoint uses the same effective bit width, group size, and symmetry, Olive defers non-canonical aliases to the canonical component build. Assembly keeps the canonical tensor and restores the tied-weight metadata. For legacy artifacts that already contain every packed alias, assembly additionally requires their tensors to be identical rather -than silently tying divergent data. If the workflow does not build the canonical component, Olive quantizes the -requested alias independently instead of deferring it; the assembled model is then untied. +than silently tying divergent data. The example uses INT8/group-32 for both ends of the tied word embedding while +quantizing the other decoder linears to INT4. If the canonical component is not built or its token table is explicitly +excluded, Olive quantizes the requested LM-head alias independently and the assembled model is untied. To keep both +tables floating point and tied when building only decoder and vision, set `lm_head: false` on the decoder pass. The named build directories contain component-only safetensors artifacts. The workflow output contains the complete checkpoint: diff --git a/olive/passes/pytorch/kquant.py b/olive/passes/pytorch/kquant.py index 20931ee68..74712f6df 100644 --- a/olive/passes/pytorch/kquant.py +++ b/olive/passes/pytorch/kquant.py @@ -289,7 +289,7 @@ class KQuant(Pass): @classmethod def _default_config(cls, accelerator_spec: AcceleratorSpec) -> dict[str, PassConfigParam]: - config = get_quantizer_config(allow_embeds=True, allow_moe=True) + config = get_quantizer_config(allow_embeds=True, allow_moe=True, auto_component_targets=True) config["group_size"] = PassConfigParam( type_=int, default_value=32, diff --git a/olive/passes/pytorch/quant_utils.py b/olive/passes/pytorch/quant_utils.py index bb3c0784f..5d32b50db 100644 --- a/olive/passes/pytorch/quant_utils.py +++ b/olive/passes/pytorch/quant_utils.py @@ -24,7 +24,7 @@ tie_quant_word_embeddings, ) from olive.common.quant.patterns import match_skip -from olive.common.quant.selection import iter_quant_targets +from olive.common.quant.selection import _collect_vision_towers, iter_quant_targets from olive.common.quant.state_dict import install_quant_tensor_param from olive.common.quant.tensor import QuantTensor from olive.common.quant.utils import WeightQuantizer @@ -46,7 +46,11 @@ _QUANTIZATION_CONFIG_NOT_PROVIDED = object() -def get_quantizer_config(allow_embeds: bool = False, allow_moe: bool = False) -> dict[str, PassConfigParam]: +def get_quantizer_config( + allow_embeds: bool = False, + allow_moe: bool = False, + auto_component_targets: bool = False, +) -> dict[str, PassConfigParam]: return { "bits": PassConfigParam( type_=PrecisionBits, @@ -68,27 +72,43 @@ def get_quantizer_config(allow_embeds: bool = False, allow_moe: bool = False) -> ), "lm_head": PassConfigParam( type_=bool, - default_value=False, + default_value=None if auto_component_targets else False, search_defaults=Boolean(), - description="Whether to quantize the language model head. Default value is False.", + description=( + "Whether to quantize the language model head. Defaults to the selected " + "component's ownership when omitted, or False for a whole-model pass." + if auto_component_targets + else "Whether to quantize the language model head. Default value is False." + ), ), "quantize_vision": PassConfigParam( type_=bool, - default_value=False, + default_value=None if auto_component_targets else False, description=( - "Whether to quantize a composite vision-language model's vision tower in this pass. When " - "False (default), the vision tower (``visual``/``vision_tower``/``vision_model``/" - "``vision_encoder``) is left in full precision -- the typical Olive pipeline quantizes it " - "separately downstream (e.g. on the ONNX side). Set to True to quantize the vision tower " - "here too, e.g. when this pass is the only quantization step for the model." + "Whether to quantize a composite vision-language model's vision tower. " + "Defaults to the selected component's ownership when omitted, or False " + "for a whole-model pass." + if auto_component_targets + else ( + "Whether to quantize a composite vision-language model's vision tower in this pass. When " + "False (default), the vision tower (``visual``/``vision_tower``/``vision_model``/" + "``vision_encoder``) is left in full precision -- the typical Olive pipeline quantizes it " + "separately downstream (e.g. on the ONNX side). Set to True to quantize the vision tower " + "here too, e.g. when this pass is the only quantization step for the model." + ) ), ), **( { "embeds": PassConfigParam( type_=bool, - default_value=False, - description="Whether to quantize the input embeddings. Default value is False.", + default_value=None if auto_component_targets else False, + description=( + "Whether to quantize input embeddings. Defaults to True for a " + "selected embedding component, or False for a whole-model pass." + if auto_component_targets + else "Whether to quantize the input embeddings. Default value is False." + ), ) } if allow_embeds @@ -364,6 +384,20 @@ def _validate_component_source_paths( ) +def _owned_vision_towers( + root_model: torch.nn.Module, + source_paths: list[str], +) -> tuple[str, ...]: + """Find vision towers intersecting the selected component paths.""" + tower_ids = {id(tower) for tower in _collect_vision_towers(root_model)} + return tuple( + name + for name, module in root_model.named_modules() + if id(module) in tower_ids + and (_is_in_component(name, source_paths) or any(_is_in_component(path, [name]) for path in source_paths)) + ) + + def _defer_shared_weight_aliases( root_model: torch.nn.Module, component_attributes: dict, @@ -377,6 +411,7 @@ def _defer_shared_weight_aliases( if not component_name: return [] workflow_components = set(component_attributes.get("workflow_components") or ()) + planned_shared_weights = component_attributes.get("workflow_planned_shared_weights") deferred = [] for shared_weight in component_attributes.get("shared_weights") or (): @@ -387,6 +422,8 @@ def _defer_shared_weight_aliases( continue if canonical["component"] not in workflow_components: continue + if planned_shared_weights is not None and shared_weight["name"] not in planned_shared_weights: + continue alias = next( (endpoint for endpoint in shared_weight.get("aliases") or () if endpoint["component"] == component_name), None, @@ -395,6 +432,8 @@ def _defer_shared_weight_aliases( continue alias_module_name = alias["parameter"].removesuffix(".weight") + if match_skip(alias_module_name, quantization.modules_to_not_convert or []): + continue if alias_module_name == lm_head_name: enabled = quantization.lm_head config_field = "lm_head" @@ -916,6 +955,25 @@ def prepare_model( if fresh_qcfg.embeds and not component_embedding_names: raise ValueError("The selected component has no torch.nn.Embedding modules to quantize.") from None + auto_targets = bool(component_name and component_name != "model" and component_source_paths) + mp_defaults = (mp_info or {}).get("default") or {} + auto_head = auto_targets and getattr(config, "lm_head", False) is None and "lm_head" not in mp_defaults + auto_embeds = auto_targets and getattr(config, "embeds", False) is None and "embeds" not in mp_defaults + auto_vision = ( + auto_targets and getattr(config, "quantize_vision", False) is None and "quantize_vision" not in mp_defaults + ) + owned_vision_towers = _owned_vision_towers(root_model, component_source_paths) if auto_targets else () + vision_tower_paths = list(owned_vision_towers) + if auto_head and component_role == "decoder" and lm_head_name is not None: + head = get_attr(root_model, lm_head_name) + fresh_qcfg.lm_head = isinstance(head, torch.nn.Linear) and _is_in_component( + lm_head_name, component_source_paths + ) + if auto_embeds and component_role == "embedding": + fresh_qcfg.embeds = bool(component_embedding_names) + if auto_vision: + fresh_qcfg.quantize_vision = bool(owned_vision_towers) + wrapper.olive_deferred_shared_weights = ( _defer_shared_weight_aliases( root_model, @@ -954,6 +1012,12 @@ def _iter_component_quant_targets( # to a common ancestor), restrict quantization to the declared sub-trees. if not _is_in_component(root_name, component_source_paths): continue + if ( + owned_vision_towers + and not quant_cfg.quantize_vision + and _is_in_component(root_name, vision_tower_paths) + ): + continue # For component-selected embedding quantization, only quantize embeddings that # belong to the selected component(s), not sibling embeddings under the slice. if pname == "weight" and isinstance(module, torch.nn.Embedding): @@ -1025,6 +1089,17 @@ def _iter_component_quant_targets( new_qargs: dict[str, dict[str, int | bool]] = { root_name: qcfg.get_qlinear_init_args(root_name) for _, _, root_name in new_targets } + if existing_qcfg is None: + if auto_head and lm_head_name not in new_qargs: + qcfg.lm_head = False + if auto_embeds and not any(name in new_qargs for name in component_embedding_names): + qcfg.embeds = False + if ( + auto_vision + and owned_vision_towers + and not any(_is_in_component(name, vision_tower_paths) for name in new_qargs) + ): + qcfg.quantize_vision = False if mp_info is not None and mp_info.get("requires_moe") is True and hasattr(config, "moe"): required_expert_targets, required_expert_overrides = _get_required_fused_expert_targets(root_model, mp_info) _validate_required_fused_expert_targets( @@ -1126,10 +1201,10 @@ def get_quant_config( "bits": config.bits, "symmetric": config.sym, "group_size": config.group_size, - "lm_head": config.lm_head, - "embeds": getattr(config, "embeds", False), + "lm_head": bool(config.lm_head), + "embeds": bool(getattr(config, "embeds", False)), "moe": getattr(config, "moe", False), - "quantize_vision": getattr(config, "quantize_vision", False), + "quantize_vision": bool(getattr(config, "quantize_vision", False)), "modules_to_not_convert": getattr(config, "modules_to_not_convert", None) or [], "overrides": deepcopy(config.overrides) if config.overrides is not None else {}, } diff --git a/olive/passes/pytorch/rtn.py b/olive/passes/pytorch/rtn.py index fbbb82655..b4b8fdbd0 100644 --- a/olive/passes/pytorch/rtn.py +++ b/olive/passes/pytorch/rtn.py @@ -27,7 +27,7 @@ class Rtn(Pass): @classmethod def _default_config(cls, accelerator_spec: AcceleratorSpec) -> dict[str, PassConfigParam]: - return get_quantizer_config(allow_embeds=True, allow_moe=True) + return get_quantizer_config(allow_embeds=True, allow_moe=True, auto_component_targets=True) @torch.no_grad() def _run_for_config( diff --git a/olive/workflows/run/builds.py b/olive/workflows/run/builds.py index 00e0507e8..0d2dbf873 100644 --- a/olive/workflows/run/builds.py +++ b/olive/workflows/run/builds.py @@ -12,6 +12,7 @@ from olive.cache import CacheConfig from olive.common.config_utils import load_config_file from olive.common.constants import DEFAULT_WORKFLOW_ID +from olive.common.quant.patterns import match_skip from olive.model import ModelConfig from olive.systems.common import SystemType from olive.workflows.run.config import BuildConfig, BuildConfigPartial, RunConfig @@ -137,6 +138,23 @@ def _paths_overlap(first: Path, second: Path) -> bool: return first == second or first in second.parents or second in first.parents +def _build_quantizes_shared_table(build_config: dict, parameter: str) -> bool: + """Whether a component build requests a tied token table from KQuant/RTN.""" + module_name = parameter.removesuffix(".weight") + for configs in build_config["passes"].values(): + for pass_config in configs if isinstance(configs, list) else [configs]: + parsed = pass_config.model_dump() if hasattr(pass_config, "model_dump") else pass_config + if parsed["type"].lower() not in {"kquant", "rtn"}: + continue + options = parsed.get("config") or parsed + if options.get("embeds") is False: + continue + if match_skip(module_name, options.get("modules_to_not_convert") or []): + continue + return True + return False + + def expand_builds(run_config: dict) -> OrderedDict[str, dict]: """Expand ``builds`` into independent, ordinary Olive run configurations.""" if not isinstance(run_config, dict): @@ -189,6 +207,25 @@ def expand_builds(run_config: dict) -> OrderedDict[str, dict]: expanded[build_name] = child_config + hf_builds = [] + for child_config in expanded.values(): + input_model = child_config.get("input_model") or {} + if input_model.get("type", "").lower() != "hfmodel": + continue + attributes = input_model["config"].get("model_attributes") or {} + if attributes.get("shared_weights"): + hf_builds.append((child_config, attributes)) + planned_shared_weights = set() + for child_config, attributes in hf_builds: + for shared_weight in attributes["shared_weights"]: + if shared_weight["kind"] != "tied_word_embeddings": + continue + if shared_weight["canonical"]["component"] != attributes.get("component_name"): + continue + if _build_quantizes_shared_table(child_config, shared_weight["canonical"]["parameter"]): + planned_shared_weights.add(shared_weight["name"]) + for _, attributes in hf_builds: + attributes["workflow_planned_shared_weights"] = sorted(planned_shared_weights) return expanded diff --git a/olive/workflows/run/hf_component_assembly.py b/olive/workflows/run/hf_component_assembly.py index c97c3fa54..1f3666357 100644 --- a/olive/workflows/run/hf_component_assembly.py +++ b/olive/workflows/run/hf_component_assembly.py @@ -272,6 +272,16 @@ def _resolve_shared_weights( endpoint_artifacts = [(endpoint, artifacts_by_component[endpoint.component]) for endpoint in endpoints] canonical_endpoint, canonical_artifact = endpoint_artifacts[0] if f"{canonical_endpoint.parameter}_qweight" not in canonical_artifact.checkpoint.keys: + if any( + request.get("name") == shared_weight.name + for _, artifact in endpoint_artifacts[1:] + for request in artifact.config.get("olive_deferred_shared_weights", ()) + ): + raise ValueError( + f"Shared weight {shared_weight.name!r} deferred an alias, but build " + f"{canonical_artifact.name!r} produced no canonical packed tensor " + f"{canonical_endpoint.parameter!r}." + ) continue endpoint_qargs = [] @@ -988,6 +998,7 @@ def try_assemble_hf_component_builds( "component_source_paths", "shared_weights", "workflow_components", + "workflow_planned_shared_weights", ): attributes.pop(name, None) attributes["assembled_components"] = [ diff --git a/test/passes/pytorch/test_kquant.py b/test/passes/pytorch/test_kquant.py index 91f142ec6..f4de1a5ee 100644 --- a/test/passes/pytorch/test_kquant.py +++ b/test/passes/pytorch/test_kquant.py @@ -18,6 +18,7 @@ from olive.passes.pytorch.kquant import KQuant, kquant_find_qparams from olive.passes.pytorch.moe_support import MoeSupportError from olive.passes.pytorch.quant_utils import prepare_model +from olive.passes.pytorch.rtn import Rtn from test.passes.pytorch.test_quantization_utils import ( DENSE_INT2_GROUP_SIZE, assert_dense_int2_mixed_precision_checkpoint, @@ -65,21 +66,7 @@ def _make_local_tiny_qwen3_moe(save_path: Path) -> HfModelHandler: def _make_local_tiny_tied_llama(save_path: Path) -> None: """Save a tiny tied model for independent component quantization.""" - from transformers import LlamaConfig, LlamaForCausalLM - - torch.manual_seed(0) - save_path.mkdir(parents=True, exist_ok=True) - config = LlamaConfig( # pylint: disable=unexpected-keyword-arg - vocab_size=32, - hidden_size=32, - intermediate_size=64, - num_hidden_layers=1, - num_attention_heads=2, - num_key_value_heads=2, - tie_word_embeddings=True, - ) - LlamaForCausalLM(config).save_pretrained(save_path) - _save_trivial_tokenizer(save_path, config.vocab_size) + make_local_tiny_dense_llama(save_path, tie_word_embeddings=True) def _is_quant(module: torch.nn.Module) -> bool: @@ -328,7 +315,6 @@ def test_kquant_defers_noncanonical_cross_component_tied_weight(tmp_path: Path): "bits": 4, "group_size": 16, "sym": True, - "lm_head": True, "overrides": {"lm_head": {"bits": 8}}, }, disable_search=True, @@ -339,7 +325,6 @@ def test_kquant_defers_noncanonical_cross_component_tied_weight(tmp_path: Path): "bits": 8, "group_size": 16, "sym": True, - "embeds": True, }, disable_search=True, ) @@ -355,6 +340,7 @@ def test_kquant_defers_noncanonical_cross_component_tied_weight(tmp_path: Path): assert not isinstance(decoder_weight, QuantTensor) assert isinstance(embedding_weight, QuantTensor) assert decoder.config.quantization_config.lm_head is False + assert embedding.config.quantization_config.embeds is True assert decoder.config.olive_deferred_shared_weights == [ { "name": "word_embeddings", @@ -417,7 +403,6 @@ def test_kquant_quantizes_alias_when_canonical_component_is_not_built( "bits": 4, "group_size": 16, "sym": True, - "lm_head": True, "overrides": {"lm_head": {"bits": 8}}, }, disable_search=True, @@ -430,9 +415,100 @@ def test_kquant_quantizes_alias_when_canonical_component_is_not_built( assert isinstance(decoder.lm_head._parameters["weight"].data, QuantTensor) assert decoder.config.quantization_config.lm_head is True + assert decoder.config.tie_word_embeddings is False + assert not isinstance(decoder.model.embed_tokens._parameters["weight"].data, QuantTensor) assert not hasattr(decoder.config, "olive_deferred_shared_weights") +@pytest.mark.parametrize("pass_type", [KQuant, Rtn]) +@pytest.mark.parametrize("embeds", [None, False]) +def test_component_embedding_selection_and_explicit_opt_out( + tmp_path: Path, + pass_type, + embeds, +): + model_path = tmp_path / "input_model" + _make_local_tiny_tied_llama(model_path) + input_model = HfModelHandler( + model_path=str(model_path), + model_attributes={ + "component_name": "embedding", + "component_role": "embedding", + "component_source_paths": ["model.embed_tokens"], + }, + ) + pass_config = {"bits": 8, "group_size": 16, "sym": True} + if embeds is not None: + pass_config["embeds"] = embeds + + quantizer = create_pass_from_dict(pass_type, pass_config, disable_search=True) + loaded = quantizer.run(input_model, str(tmp_path / "output")).load_model() + + assert isinstance(loaded.model.embed_tokens._parameters["weight"].data, QuantTensor) is (embeds is None) + assert loaded.config.quantization_config.embeds is (embeds is None) + assert not isinstance(loaded.lm_head._parameters["weight"].data, QuantTensor) + + +@pytest.mark.parametrize( + "head_options", + [ + {"lm_head": False}, + {"modules_to_not_convert": ["lm_head"]}, + ], +) +def test_component_decoder_explicit_lm_head_opt_out(tmp_path: Path, head_options): + model_path = tmp_path / "input_model" + _make_local_tiny_tied_llama(model_path) + input_model = HfModelHandler( + model_path=str(model_path), + model_attributes={ + "component_name": "decoder", + "component_role": "decoder", + "component_source_paths": ["model.layers", "model.norm", "lm_head"], + }, + ) + quantizer = create_pass_from_dict( + KQuant, + { + "bits": 4, + "group_size": 16, + "sym": True, + "overrides": {"lm_head": {"bits": 8}}, + **head_options, + }, + disable_search=True, + ) + + loaded = quantizer.run(input_model, str(tmp_path / "output")).load_model() + + assert not isinstance(loaded.lm_head._parameters["weight"].data, QuantTensor) + assert loaded.config.quantization_config.lm_head is False + + +@pytest.mark.parametrize("pass_type", [KQuant, Rtn]) +def test_whole_model_omitted_flags_still_leave_tied_tables_float( + tmp_path: Path, + pass_type, +): + model_path = tmp_path / "input_model" + _make_local_tiny_tied_llama(model_path) + quantizer = create_pass_from_dict( + pass_type, + {"bits": 4, "group_size": 16, "sym": True}, + disable_search=True, + ) + + loaded = quantizer.run( + HfModelHandler(model_path=str(model_path)), + str(tmp_path / "output"), + ).load_model() + + assert not isinstance(loaded.lm_head._parameters["weight"].data, QuantTensor) + assert not isinstance(loaded.model.embed_tokens._parameters["weight"].data, QuantTensor) + assert loaded.config.quantization_config.lm_head is False + assert loaded.config.quantization_config.embeds is False + + @pytest.mark.parametrize( ("layout", "message"), [ diff --git a/test/passes/pytorch/test_quant_utils.py b/test/passes/pytorch/test_quant_utils.py index 75f03c469..45879cb10 100644 --- a/test/passes/pytorch/test_quant_utils.py +++ b/test/passes/pytorch/test_quant_utils.py @@ -1070,6 +1070,33 @@ def test_prepare_model_component_generated_exclusions_are_exact(input_model, mon assert not match_skip("blocks.10", qcfg.modules_to_not_convert) +@pytest.mark.parametrize("quantize_vision", [None, False]) +def test_prepare_model_scoped_vision_auto_selection_and_opt_out(input_model, monkeypatch, quantize_vision): + root_model = _make_nested_decoder_root(input_model) + root_model.config.vision_config = SimpleNamespace() + root_model.vision_tower = torch.nn.Linear(16, 16) + monkeypatch.setattr(quant_utils_module, "load_hf_base_model", lambda _: root_model) + model = HfModelHandler( + input_model.model_path, + model_attributes={ + "component_name": "vision_encoder", + "component_role": "encoder", + "component_source_paths": ["vision_tower"], + }, + ) + config = _baseline_pass_config() + config.quantize_vision = quantize_vision + + _, qcfg, _ = prepare_model(model, config) + + assert qcfg.quantize_vision is (quantize_vision is None) + assert hasattr(root_model.vision_tower.weight, "quant_info") is (quantize_vision is None) + assert not hasattr( + root_model.decoder.model.layers[0].self_attn.q_proj.weight, + "quant_info", + ) + + def test_finalize_multi_path_vlm_decoder_quantizes_and_saves_full_model( input_model, monkeypatch, diff --git a/test/passes/pytorch/test_quantization_utils.py b/test/passes/pytorch/test_quantization_utils.py index e02aff803..e42d32e99 100644 --- a/test/passes/pytorch/test_quantization_utils.py +++ b/test/passes/pytorch/test_quantization_utils.py @@ -32,7 +32,7 @@ def _dense_int2_calibration_dataset(seq_len: int, max_samples: int, vocab_size: ] -def make_local_tiny_dense_llama(save_path: Path) -> HfModelHandler: +def make_local_tiny_dense_llama(save_path: Path, *, tie_word_embeddings: bool = False) -> HfModelHandler: """Save a tiny dense Llama checkpoint and tokenizer without accessing the hub.""" from tokenizers import Tokenizer, models, pre_tokenizers from transformers import LlamaConfig, LlamaForCausalLM, PreTrainedTokenizerFast @@ -46,6 +46,7 @@ def make_local_tiny_dense_llama(save_path: Path) -> HfModelHandler: num_hidden_layers=1, num_attention_heads=2, num_key_value_heads=2, + **({"tie_word_embeddings": True} if tie_word_embeddings else {}), ) LlamaForCausalLM(config).save_pretrained(save_path) diff --git a/test/workflows/test_hf_component_assembly.py b/test/workflows/test_hf_component_assembly.py index acddfb5d3..55ebc7308 100644 --- a/test/workflows/test_hf_component_assembly.py +++ b/test/workflows/test_hf_component_assembly.py @@ -76,6 +76,8 @@ def _run_config( pass_type: str, *, shared_weights: list[dict] | None = None, + role: str | None = None, + workflow_components: list[str] | None = None, ): source_paths = [source_path] if isinstance(source_path, str) else source_path model_attributes = { @@ -84,6 +86,10 @@ def _run_config( } if shared_weights: model_attributes["shared_weights"] = shared_weights + if role is not None: + model_attributes["component_role"] = role + if workflow_components is not None: + model_attributes["workflow_components"] = workflow_components return SimpleNamespace( input_model=SimpleNamespace( type="hfmodel", @@ -299,6 +305,7 @@ def test_assembles_disjoint_hf_components_and_preserves_unbuilt_weights(tmp_path [ (None, None), ("deferred", None), + ("missing_canonical", "produced no canonical packed tensor"), ("layout", "incompatible quantization layouts"), ("tensor", "differs from canonical tensor"), ], @@ -335,8 +342,8 @@ def test_assembles_cross_component_tied_word_embeddings(tmp_path, failure, messa symmetric=True, quantize_vision=False, skips=[], - overrides=None if failure == "deferred" else {"lm_head": {"bits": 8}}, - lm_head=failure != "deferred", + overrides=(None if failure in {"deferred", "missing_canonical"} else {"lm_head": {"bits": 8}}), + lm_head=failure not in {"deferred", "missing_canonical"}, ) embedding_quantization = _quantization_config( group_size=32, @@ -349,7 +356,7 @@ def test_assembles_cross_component_tied_word_embeddings(tmp_path, failure, messa **({"group_size": 16} if failure == "layout" else {}), } }, - embeds=True, + embeds=failure != "missing_canonical", ) tied_qweight = torch.arange(64, dtype=torch.uint8).reshape(8, 8) embedding_qweight = tied_qweight.clone() @@ -361,7 +368,7 @@ def test_assembles_cross_component_tied_word_embeddings(tmp_path, failure, messa "text_config": {"tie_word_embeddings": False}, } decoder_model_config = dict(model_config) - if failure == "deferred": + if failure in {"deferred", "missing_canonical"}: decoder_model_config["olive_deferred_shared_weights"] = [ { "name": "word_embeddings", @@ -380,7 +387,7 @@ def test_assembles_cross_component_tied_word_embeddings(tmp_path, failure, messa "model.language_model.layers.0.weight_scales": torch.ones(8, 2), "model.unoptimized.weight": torch.ones(8, 8), } - if failure != "deferred": + if failure not in {"deferred", "missing_canonical"}: decoder_tensors.update( { "lm_head.weight_qweight": tied_qweight, @@ -393,13 +400,19 @@ def test_assembles_cross_component_tied_word_embeddings(tmp_path, failure, messa decoder_quantization, model_config=decoder_model_config, ) + embedding_tensors = {"model.unoptimized.weight": torch.ones(8, 8)} + if failure == "missing_canonical": + embedding_tensors["model.language_model.embed_tokens.weight"] = torch.ones(8, 8) + else: + embedding_tensors.update( + { + "model.language_model.embed_tokens.weight_qweight": embedding_qweight, + "model.language_model.embed_tokens.weight_scales": tied_scales, + } + ) _write_checkpoint( embedding_model, - { - "model.language_model.embed_tokens.weight_qweight": embedding_qweight, - "model.language_model.embed_tokens.weight_scales": tied_scales, - "model.unoptimized.weight": torch.ones(8, 8), - }, + embedding_tensors, embedding_quantization, model_config=model_config, ) @@ -482,6 +495,91 @@ def test_assembles_cross_component_tied_word_embeddings(tmp_path, failure, messa assert embedding_manifest["quantization_config"]["tie_word_embeddings"] is True +@pytest.mark.parametrize("embedding_enabled", [True, False]) +def test_assembles_auto_selected_tied_component_quantization(tmp_path, embedding_enabled): + from olive.common.quant.tensor import QuantTensor + from olive.model import HfModelHandler + from olive.passes.olive_pass import create_pass_from_dict + from olive.passes.pytorch.kquant import KQuant + from test.passes.pytorch.test_quantization_utils import make_local_tiny_dense_llama + + source = tmp_path / "source" + make_local_tiny_dense_llama(source, tie_word_embeddings=True) + shared_weights = [ + { + "name": "word_embeddings", + "kind": "tied_word_embeddings", + "canonical": { + "component": "embedding", + "parameter": "model.embed_tokens.weight", + }, + "aliases": [{"component": "decoder", "parameter": "lm_head.weight"}], + } + ] + components = ["decoder", "embedding"] + paths = { + "decoder": ["model.layers", "model.norm", "lm_head"], + "embedding": ["model.embed_tokens"], + } + pass_configs = { + "decoder": { + "bits": 4, + "group_size": 16, + "sym": True, + "overrides": {"lm_head": {"bits": 8}}, + }, + "embedding": {"bits": 8, "group_size": 16, "sym": True}, + } + if not embedding_enabled: + pass_configs["embedding"]["embeds"] = False + planned_shared_weights = ["word_embeddings"] if embedding_enabled else [] + build_configs = OrderedDict() + results = OrderedDict() + for component in components: + output_dir = tmp_path / "builds" / component + model_dir = output_dir / "model" + input_model = HfModelHandler( + model_path=str(source), + model_attributes={ + "component_name": component, + "component_role": component, + "component_source_paths": paths[component], + "shared_weights": shared_weights, + "workflow_components": components, + "workflow_planned_shared_weights": planned_shared_weights, + }, + ) + quantizer = create_pass_from_dict(KQuant, pass_configs[component], disable_search=True) + quantizer.run(input_model, str(model_dir)) + build_configs[component] = _run_config( + output_dir, + component, + paths[component], + "KQuant", + shared_weights=shared_weights, + role=component, + workflow_components=components, + ) + results[component] = _result(model_dir) + + assembled_dir = tmp_path / "assembled" + assert try_assemble_hf_component_builds(build_configs, results, assembled_dir) == assembled_dir + + quantization = json.loads((assembled_dir / "config.json").read_text())["quantization_config"] + assert quantization["lm_head"] is True + assert quantization["embeds"] is embedding_enabled + assert quantization["tie_word_embeddings"] is embedding_enabled + keys = _checkpoint_keys(assembled_dir) + assert ("model.embed_tokens.weight_qweight" in keys) is embedding_enabled + assert ("lm_head.weight_qweight" in keys) is not embedding_enabled + loaded = HfModelHandler(model_path=str(assembled_dir)).load_model() + embedding = loaded.get_input_embeddings()._parameters["weight"] + head = loaded.get_output_embeddings()._parameters["weight"] + assert (embedding is head) is embedding_enabled + assert isinstance(embedding.data, QuantTensor) is embedding_enabled + assert isinstance(head.data, QuantTensor) + + def test_quantization_merge_resolves_effective_overrides_and_float_skips(): decoder_config = _quantization_config( group_size=32, diff --git a/test/workflows/test_run_builds.py b/test/workflows/test_run_builds.py index 972cbeced..71e885f61 100644 --- a/test/workflows/test_run_builds.py +++ b/test/workflows/test_run_builds.py @@ -3,6 +3,7 @@ # Licensed under the MIT License. # -------------------------------------------------------------------------- +import json import sys from copy import deepcopy from pathlib import Path @@ -90,6 +91,156 @@ def test_builds_tag_all_selected_hf_components(self, monkeypatch): attributes = build["input_model"]["config"]["model_attributes"] assert attributes["workflow_components"] == ["decoder", "embedding"] + @pytest.mark.parametrize( + ("embedding_options", "planned"), + [ + ({}, ["word_embeddings"]), + ({"embeds": False}, []), + ({"modules_to_not_convert": ["model.embed_tokens"]}, []), + ], + ) + def test_builds_plan_only_quantized_shared_owners(self, monkeypatch, embedding_options, planned): + shared_weight = mobius_utils.SharedWeightInfo.coerce( + { + "name": "word_embeddings", + "kind": "tied_word_embeddings", + "canonical": { + "component": "embedding", + "parameter": "model.embed_tokens.weight", + }, + "aliases": [{"component": "decoder", "parameter": "lm_head.weight"}], + } + ) + monkeypatch.setattr( + mobius_utils, + "inspect_components", + lambda *args, **kwargs: [ + mobius_utils.ComponentInfo( + name="decoder", + role="decoder", + source_paths=["model.layers", "lm_head"], + shared_weights=[shared_weight], + ), + mobius_utils.ComponentInfo( + name="embedding", + role="embedding", + source_paths=["model.embed_tokens"], + shared_weights=[shared_weight], + ), + ], + ) + config = deepcopy(self.template) + config["input_model"] = { + "type": "HfModel", + "config": {"model_path": "local/model"}, + } + config["passes"] = { + "decoder_quant": {"type": "KQuant"}, + "embedding_quant": {"type": "KQuant", **embedding_options}, + } + config["builds"] = { + "decoder": { + "components": ["decoder"], + "pipeline": ["decoder_quant"], + }, + "embedding": { + "components": ["embedding"], + "pipeline": ["embedding_quant"], + }, + } + + expanded = expand_builds(config) + + for build in expanded.values(): + attributes = build["input_model"]["config"]["model_attributes"] + assert attributes["workflow_components"] == ["decoder", "embedding"] + assert attributes["workflow_planned_shared_weights"] == planned + + def test_builds_assemble_auto_selected_tied_weights(self, monkeypatch, tmp_path): + from olive.common.quant.tensor import QuantTensor + from olive.model import HfModelHandler + from test.passes.pytorch.test_quantization_utils import make_local_tiny_dense_llama + + source = tmp_path / "source" + make_local_tiny_dense_llama(source, tie_word_embeddings=True) + shared_weight = mobius_utils.SharedWeightInfo.coerce( + { + "name": "word_embeddings", + "kind": "tied_word_embeddings", + "canonical": { + "component": "embedding", + "parameter": "model.embed_tokens.weight", + }, + "aliases": [{"component": "decoder", "parameter": "lm_head.weight"}], + } + ) + monkeypatch.setattr( + mobius_utils, + "inspect_components", + lambda *args, **kwargs: [ + mobius_utils.ComponentInfo( + name="decoder", + role="decoder", + source_paths=["model.layers", "model.norm", "lm_head"], + shared_weights=[shared_weight], + ), + mobius_utils.ComponentInfo( + name="embedding", + role="embedding", + source_paths=["model.embed_tokens"], + shared_weights=[shared_weight], + ), + ], + ) + output_dir = tmp_path / "assembled" + config = { + "input_model": {"type": "HfModel", "config": {"model_path": str(source)}}, + "passes": { + "decoder_kquant": { + "type": "KQuant", + "bits": 4, + "group_size": 16, + "sym": True, + "overrides": {"lm_head": {"bits": 8}}, + }, + "embedding_kquant": { + "type": "KQuant", + "bits": 8, + "group_size": 16, + "sym": True, + }, + }, + "builds": { + "decoder": { + "components": ["decoder"], + "pipeline": ["decoder_kquant"], + }, + "embedding": { + "components": ["embedding"], + "pipeline": ["embedding_kquant"], + }, + }, + "engine": { + "output_dir": str(output_dir), + "cache_dir": str(tmp_path / "cache"), + "evaluate_input_model": False, + }, + "max_concurrent_builds": 1, + } + + olive_run(config) + + checkpoint = json.loads((output_dir / "model.safetensors.index.json").read_text(encoding="utf-8")) + assert "model.embed_tokens.weight_qweight" in checkpoint["weight_map"] + assert "lm_head.weight_qweight" not in checkpoint["weight_map"] + loaded = HfModelHandler(model_path=str(output_dir)).load_model() + assert loaded.config.quantization_config.lm_head is True + assert loaded.config.quantization_config.embeds is True + assert loaded.config.quantization_config.tie_word_embeddings is True + embedding = loaded.get_input_embeddings()._parameters["weight"] + assert embedding is loaded.get_output_embeddings()._parameters["weight"] + assert isinstance(embedding.data, QuantTensor) + def test_builds_components_on_non_composite_input_raises(self): config = deepcopy(self.template) config["builds"] = { From 5874daa91e017fc760356d818ae8ce39d7423df9 Mon Sep 17 00:00:00 2001 From: Xiaoyu Zhang Date: Thu, 24 Sep 2026 10:56:44 -0700 Subject: [PATCH 5/9] Reuse tied-weight fixtures in scoped build tests Share the synthetic tied-word-embedding declaration across quantization and workflow tests, and keep the new assertions compatible with the repository's Pylint configuration. Signed-off-by: Xiaoyu Zhang --- test/passes/pytorch/test_kquant.py | 35 ++----------------- test/passes/pytorch/test_quant_utils.py | 5 +-- .../passes/pytorch/test_quantization_utils.py | 13 +++++++ test/workflows/test_hf_component_assembly.py | 30 ++-------------- test/workflows/test_run_builds.py | 29 +++------------ 5 files changed, 27 insertions(+), 85 deletions(-) diff --git a/test/passes/pytorch/test_kquant.py b/test/passes/pytorch/test_kquant.py index f4de1a5ee..7f15464c1 100644 --- a/test/passes/pytorch/test_kquant.py +++ b/test/passes/pytorch/test_kquant.py @@ -25,6 +25,7 @@ assert_uniform_int2_checkpoint, make_local_tiny_dense_llama, plan_dense_int2_mixed_precision, + tied_word_embedding_group, ) from test.utils import get_tiny_phi3 @@ -268,22 +269,7 @@ def test_kquant_consumes_selective_mixed_precision_int2_int4_int8(tmp_path: Path def test_kquant_defers_noncanonical_cross_component_tied_weight(tmp_path: Path): model_path = tmp_path / "input_model" _make_local_tiny_tied_llama(model_path) - shared_weights = [ - { - "name": "word_embeddings", - "kind": "tied_word_embeddings", - "canonical": { - "component": "embedding", - "parameter": "model.embed_tokens.weight", - }, - "aliases": [ - { - "component": "decoder", - "parameter": "lm_head.weight", - } - ], - } - ] + shared_weights = [tied_word_embedding_group()] decoder_model = HfModelHandler( model_path=str(model_path), model_attributes={ @@ -379,22 +365,7 @@ def test_kquant_quantizes_alias_when_canonical_component_is_not_built( "lm_head", ], "workflow_components": ["decoder", "vision_encoder"], - "shared_weights": [ - { - "name": "word_embeddings", - "kind": "tied_word_embeddings", - "canonical": { - "component": "embedding", - "parameter": "model.embed_tokens.weight", - }, - "aliases": [ - { - "component": "decoder", - "parameter": "lm_head.weight", - } - ], - } - ], + "shared_weights": [tied_word_embedding_group()], }, ) decoder_pass = create_pass_from_dict( diff --git a/test/passes/pytorch/test_quant_utils.py b/test/passes/pytorch/test_quant_utils.py index 45879cb10..56aa5ce70 100644 --- a/test/passes/pytorch/test_quant_utils.py +++ b/test/passes/pytorch/test_quant_utils.py @@ -1074,7 +1074,8 @@ def test_prepare_model_component_generated_exclusions_are_exact(input_model, mon def test_prepare_model_scoped_vision_auto_selection_and_opt_out(input_model, monkeypatch, quantize_vision): root_model = _make_nested_decoder_root(input_model) root_model.config.vision_config = SimpleNamespace() - root_model.vision_tower = torch.nn.Linear(16, 16) + vision_tower = torch.nn.Linear(16, 16) + root_model.add_module("vision_tower", vision_tower) monkeypatch.setattr(quant_utils_module, "load_hf_base_model", lambda _: root_model) model = HfModelHandler( input_model.model_path, @@ -1090,7 +1091,7 @@ def test_prepare_model_scoped_vision_auto_selection_and_opt_out(input_model, mon _, qcfg, _ = prepare_model(model, config) assert qcfg.quantize_vision is (quantize_vision is None) - assert hasattr(root_model.vision_tower.weight, "quant_info") is (quantize_vision is None) + assert hasattr(vision_tower.weight, "quant_info") is (quantize_vision is None) assert not hasattr( root_model.decoder.model.layers[0].self_attn.q_proj.weight, "quant_info", diff --git a/test/passes/pytorch/test_quantization_utils.py b/test/passes/pytorch/test_quantization_utils.py index e42d32e99..53ae64df3 100644 --- a/test/passes/pytorch/test_quantization_utils.py +++ b/test/passes/pytorch/test_quantization_utils.py @@ -56,6 +56,19 @@ def make_local_tiny_dense_llama(save_path: Path, *, tie_word_embeddings: bool = return HfModelHandler(model_path=str(save_path)) +def tied_word_embedding_group(embedding_parameter: str = "model.embed_tokens.weight") -> dict: + """Describe a tied token table owned by a separate embedding component.""" + return { + "name": "word_embeddings", + "kind": "tied_word_embeddings", + "canonical": { + "component": "embedding", + "parameter": embedding_parameter, + }, + "aliases": [{"component": "decoder", "parameter": "lm_head.weight"}], + } + + def make_local_calibration_data_config(seq_len: int = 16, max_samples: int = 4) -> DataConfig: """Return deterministic token calibration data generated entirely in-process.""" return DataConfig( diff --git a/test/workflows/test_hf_component_assembly.py b/test/workflows/test_hf_component_assembly.py index 55ebc7308..c55f4f01d 100644 --- a/test/workflows/test_hf_component_assembly.py +++ b/test/workflows/test_hf_component_assembly.py @@ -23,6 +23,7 @@ _validate_build_compatibility, try_assemble_hf_component_builds, ) +from test.passes.pytorch.test_quantization_utils import tied_word_embedding_group def _quantization_config( @@ -321,22 +322,7 @@ def test_assembles_cross_component_tied_word_embeddings(tmp_path, failure, messa (decoder_output / "model_config.json").write_text("{}", encoding="utf-8") (embedding_output / "model_config.json").write_text("{}", encoding="utf-8") - shared_weights = [ - { - "name": "word_embeddings", - "kind": "tied_word_embeddings", - "canonical": { - "component": "embedding", - "parameter": "model.language_model.embed_tokens.weight", - }, - "aliases": [ - { - "component": "decoder", - "parameter": "lm_head.weight", - } - ], - } - ] + shared_weights = [tied_word_embedding_group("model.language_model.embed_tokens.weight")] decoder_quantization = _quantization_config( group_size=32, symmetric=True, @@ -505,17 +491,7 @@ def test_assembles_auto_selected_tied_component_quantization(tmp_path, embedding source = tmp_path / "source" make_local_tiny_dense_llama(source, tie_word_embeddings=True) - shared_weights = [ - { - "name": "word_embeddings", - "kind": "tied_word_embeddings", - "canonical": { - "component": "embedding", - "parameter": "model.embed_tokens.weight", - }, - "aliases": [{"component": "decoder", "parameter": "lm_head.weight"}], - } - ] + shared_weights = [tied_word_embedding_group()] components = ["decoder", "embedding"] paths = { "decoder": ["model.layers", "model.norm", "lm_head"], diff --git a/test/workflows/test_run_builds.py b/test/workflows/test_run_builds.py index 71e885f61..9a638fac3 100644 --- a/test/workflows/test_run_builds.py +++ b/test/workflows/test_run_builds.py @@ -16,6 +16,7 @@ from olive.common import mobius_utils from olive.workflows import run as olive_run from olive.workflows.run.builds import expand_builds +from test.passes.pytorch.test_quantization_utils import tied_word_embedding_group from test.utils import get_pytorch_model_io_config, pytorch_model_loader # pylint: disable=attribute-defined-outside-init @@ -100,17 +101,7 @@ def test_builds_tag_all_selected_hf_components(self, monkeypatch): ], ) def test_builds_plan_only_quantized_shared_owners(self, monkeypatch, embedding_options, planned): - shared_weight = mobius_utils.SharedWeightInfo.coerce( - { - "name": "word_embeddings", - "kind": "tied_word_embeddings", - "canonical": { - "component": "embedding", - "parameter": "model.embed_tokens.weight", - }, - "aliases": [{"component": "decoder", "parameter": "lm_head.weight"}], - } - ) + shared_weight = mobius_utils.SharedWeightInfo.coerce(tied_word_embedding_group()) monkeypatch.setattr( mobius_utils, "inspect_components", @@ -163,17 +154,7 @@ def test_builds_assemble_auto_selected_tied_weights(self, monkeypatch, tmp_path) source = tmp_path / "source" make_local_tiny_dense_llama(source, tie_word_embeddings=True) - shared_weight = mobius_utils.SharedWeightInfo.coerce( - { - "name": "word_embeddings", - "kind": "tied_word_embeddings", - "canonical": { - "component": "embedding", - "parameter": "model.embed_tokens.weight", - }, - "aliases": [{"component": "decoder", "parameter": "lm_head.weight"}], - } - ) + shared_weight = mobius_utils.SharedWeightInfo.coerce(tied_word_embedding_group()) monkeypatch.setattr( mobius_utils, "inspect_components", @@ -237,8 +218,8 @@ def test_builds_assemble_auto_selected_tied_weights(self, monkeypatch, tmp_path) assert loaded.config.quantization_config.lm_head is True assert loaded.config.quantization_config.embeds is True assert loaded.config.quantization_config.tie_word_embeddings is True - embedding = loaded.get_input_embeddings()._parameters["weight"] - assert embedding is loaded.get_output_embeddings()._parameters["weight"] + embedding = loaded.get_input_embeddings().weight + assert embedding is loaded.get_output_embeddings().weight assert isinstance(embedding.data, QuantTensor) def test_builds_components_on_non_composite_input_raises(self): From 00ab65908a4342a54d7941419fa13f5fd7272ea2 Mon Sep 17 00:00:00 2001 From: Xiaoyu Zhang Date: Thu, 24 Sep 2026 11:35:07 -0700 Subject: [PATCH 6/9] Preserve unscoped HF builds and ComponentInfo constructor order Run shared-weight planning only for explicitly selected component builds so ordinary flattened HfModel configs remain valid. Restore metadata as ComponentInfo's fourth positional argument and cover both compatibility paths with regression tests. Signed-off-by: Xiaoyu Zhang --- olive/common/mobius_utils.py | 4 +++- olive/workflows/run/builds.py | 8 +++++--- test/common/test_mobius_utils.py | 9 +++++++++ test/workflows/test_run_config_builds.py | 10 ++++++++++ 4 files changed, 27 insertions(+), 4 deletions(-) diff --git a/olive/common/mobius_utils.py b/olive/common/mobius_utils.py index 3a9a4a586..e5eb582cf 100644 --- a/olive/common/mobius_utils.py +++ b/olive/common/mobius_utils.py @@ -105,14 +105,16 @@ class ComponentInfo: source_paths: Dotted submodule paths locating the component inside the full model (e.g. ``["model.language_model"]``). A component may span multiple disjoint sub-modules, so this is a list. + metadata: Additional component metadata retained from earlier callers. + shared_weights: Cross-component shared-weight declarations from Mobius. """ name: str role: Optional[str] = None source_paths: list[str] = field(default_factory=list) - shared_weights: list[SharedWeightInfo] = field(default_factory=list) metadata: dict = field(default_factory=dict) + shared_weights: list[SharedWeightInfo] = field(default_factory=list) @classmethod def coerce(cls, data: "ComponentInfo | dict | object") -> "ComponentInfo": diff --git a/olive/workflows/run/builds.py b/olive/workflows/run/builds.py index 0d2dbf873..2e1cc0398 100644 --- a/olive/workflows/run/builds.py +++ b/olive/workflows/run/builds.py @@ -208,9 +208,11 @@ def expand_builds(run_config: dict) -> OrderedDict[str, dict]: expanded[build_name] = child_config hf_builds = [] - for child_config in expanded.values(): - input_model = child_config.get("input_model") or {} - if input_model.get("type", "").lower() != "hfmodel": + for build_name, child_config in expanded.items(): + if not builds[build_name].components: + continue + input_model = child_config["input_model"] + if input_model["type"].lower() != "hfmodel": continue attributes = input_model["config"].get("model_attributes") or {} if attributes.get("shared_weights"): diff --git a/test/common/test_mobius_utils.py b/test/common/test_mobius_utils.py index 8f525c592..708195701 100644 --- a/test/common/test_mobius_utils.py +++ b/test/common/test_mobius_utils.py @@ -22,6 +22,15 @@ def test_coerce_reads_contract_dict(): assert component.metadata == {"extra": 1} +def test_component_info_keeps_metadata_as_fourth_positional_argument(): + metadata = {"source": "legacy"} + + component = ComponentInfo("decoder", "decoder", ["model.layers"], metadata) + + assert component.metadata is metadata + assert not component.shared_weights + + def test_coerce_reads_mobius_source_paths_tuple(): # A component may span multiple disjoint HF sub-trees (e.g. phi4mm decoder). component = ComponentInfo.coerce( diff --git a/test/workflows/test_run_config_builds.py b/test/workflows/test_run_config_builds.py index acc135cc7..cb01a8ead 100644 --- a/test/workflows/test_run_config_builds.py +++ b/test/workflows/test_run_config_builds.py @@ -61,6 +61,16 @@ def test_ordinary_run_config_with_null_builds_round_trips(self): assert "max_concurrent_builds" not in serialized assert isinstance(parse_run_config(serialized), RunConfig) + def test_unscoped_hf_builds_keep_flat_input_model(self): + builds = self._expand( + { + "first": {"pipeline": ["convert"]}, + "second": {"pipeline": ["tune"]}, + } + ) + + assert all(build["input_model"] == self.template["input_model"] for build in builds.values()) + def test_builds_prevalidate_duplicate_output_dirs(self): config = deepcopy(self.template) config["builds"] = { From 4112c856b2e3e7f35e3e82deb45703a65f698820 Mon Sep 17 00:00:00 2001 From: Xiaoyu Zhang Date: Thu, 24 Sep 2026 11:44:47 -0700 Subject: [PATCH 7/9] Keep component build documentation concise Remove redundant quantization behavior explanations and show the embedding artifact in the three-build example output. Signed-off-by: Xiaoyu Zhang --- .../configure-workflows/build-workflow.md | 18 ++---------------- 1 file changed, 2 insertions(+), 16 deletions(-) diff --git a/docs/source/how-to/configure-workflows/build-workflow.md b/docs/source/how-to/configure-workflows/build-workflow.md index ab9faee6a..14dfe867c 100644 --- a/docs/source/how-to/configure-workflows/build-workflow.md +++ b/docs/source/how-to/configure-workflows/build-workflow.md @@ -188,22 +188,6 @@ By default, each named build is saved under `/`. to any other location without changing where the assembled model is saved. Olive refuses to assemble into a workflow output directory that already contains files. -For KQuant and RTN on one selected Hugging Face component, omitted `lm_head`, `embeds`, and `quantize_vision` -automatically include the selected decoder's LM head, an embedding component's token tables, and owned vision-tower -weights, respectively. Explicit `false` keeps a category floating point; `modules_to_not_convert` excludes -individual modules. Whole-model passes retain their original opt-in defaults. The saved component and assembled -`quantization_config` still record `lm_head` and `embeds` as the actual checkpoint layout, including deferred -cross-component aliases. - -When Mobius reports that parameters are tied across components, keep each component in its own build. If every -requested endpoint uses the same effective bit width, group size, and symmetry, Olive defers non-canonical aliases to -the canonical component build. Assembly keeps the canonical tensor and restores the tied-weight metadata. For legacy -artifacts that already contain every packed alias, assembly additionally requires their tensors to be identical rather -than silently tying divergent data. The example uses INT8/group-32 for both ends of the tied word embedding while -quantizing the other decoder linears to INT4. If the canonical component is not built or its token table is explicitly -excluded, Olive quantizes the requested LM-head alias independently and the assembled model is untied. To keep both -tables floating point and tied when building only decoder and vision, set `lm_head: false` on the decoder pass. - The named build directories contain component-only safetensors artifacts. The workflow output contains the complete checkpoint: @@ -214,6 +198,8 @@ models/gemma4/ model-unoptimized-00001.safetensors decoder/model-00001.safetensors decoder/component.json + embedding/model-00001.safetensors + embedding/component.json vision/model-00001.safetensors vision/component.json ``` From 9c5534e044d3f5c60853253359ba27f4dfa75685 Mon Sep 17 00:00:00 2001 From: Xiaoyu Zhang Date: Fri, 25 Sep 2026 12:33:34 -0700 Subject: [PATCH 8/9] Make cross-component tied quantization fail closed Preserve deferred LM heads across follow-up passes and require HF assembly when shared weights are pending. Use the pass's effective mixed-precision layout to validate both endpoints before building, recover float output heads from the original tied table for one-sided quantization, and reject malformed or unsupported sharing metadata. Signed-off-by: Xiaoyu Zhang --- olive/common/mobius_utils.py | 32 +++- olive/passes/pytorch/quant_utils.py | 45 ++++-- olive/workflows/run/builds.py | 58 +++++-- olive/workflows/run/hf_component_assembly.py | 60 ++++++- olive/workflows/run/run.py | 11 +- test/common/test_mobius_utils.py | 28 ++++ test/passes/pytorch/test_kquant.py | 36 +++++ test/workflows/test_hf_component_assembly.py | 112 +++++++++++-- test/workflows/test_run_builds.py | 158 ++++++++++++------- 9 files changed, 445 insertions(+), 95 deletions(-) diff --git a/olive/common/mobius_utils.py b/olive/common/mobius_utils.py index e5eb582cf..ef0e35d84 100644 --- a/olive/common/mobius_utils.py +++ b/olive/common/mobius_utils.py @@ -35,16 +35,22 @@ class SharedWeightEndpoint: component: str parameter: str + def __post_init__(self) -> None: + if not isinstance(self.component, str) or not self.component: + raise ValueError("Shared-weight endpoint component must be a non-empty string.") + if not isinstance(self.parameter, str) or not self.parameter.endswith(".weight"): + raise ValueError("Shared-weight endpoint parameter must name a '.weight' tensor.") + @classmethod def coerce(cls, data: "SharedWeightEndpoint | dict | object") -> "SharedWeightEndpoint": if isinstance(data, cls): return data if isinstance(data, dict): - return cls(component=str(data["component"]), parameter=str(data["parameter"])) + return cls(component=data["component"], parameter=data["parameter"]) duck_data: Any = data return cls( - component=str(duck_data.component), - parameter=str(duck_data.parameter), + component=duck_data.component, + parameter=duck_data.parameter, ) def to_json(self) -> dict[str, str]: @@ -60,23 +66,35 @@ class SharedWeightInfo: aliases: list[SharedWeightEndpoint] = field(default_factory=list) kind: str = "parameter_alias" + def __post_init__(self) -> None: + if not isinstance(self.name, str) or not self.name: + raise ValueError("Shared-weight name must be a non-empty string.") + if not isinstance(self.kind, str) or not self.kind: + raise ValueError(f"Shared weight {self.name!r} must declare a kind.") + if not self.aliases: + raise ValueError(f"Shared weight {self.name!r} must declare at least one alias.") + endpoints = self.endpoints + names = [endpoint.parameter for endpoint in endpoints] + if len(set(names)) != len(names): + raise ValueError(f"Shared weight {self.name!r} contains duplicate parameter endpoints.") + @classmethod def coerce(cls, data: "SharedWeightInfo | dict | object") -> "SharedWeightInfo": if isinstance(data, cls): return data if isinstance(data, dict): return cls( - name=str(data["name"]), + name=data["name"], canonical=SharedWeightEndpoint.coerce(data["canonical"]), aliases=[SharedWeightEndpoint.coerce(alias) for alias in data.get("aliases", ())], - kind=str(data.get("kind", "parameter_alias")), + kind=data.get("kind", "parameter_alias"), ) duck_data: Any = data return cls( - name=str(duck_data.name), + name=duck_data.name, canonical=SharedWeightEndpoint.coerce(duck_data.canonical), aliases=[SharedWeightEndpoint.coerce(alias) for alias in duck_data.aliases], - kind=str(getattr(duck_data, "kind", "parameter_alias")), + kind=getattr(duck_data, "kind", "parameter_alias"), ) @property diff --git a/olive/passes/pytorch/quant_utils.py b/olive/passes/pytorch/quant_utils.py index 5d32b50db..689f975f4 100644 --- a/olive/passes/pytorch/quant_utils.py +++ b/olive/passes/pytorch/quant_utils.py @@ -412,6 +412,7 @@ def _defer_shared_weight_aliases( return [] workflow_components = set(component_attributes.get("workflow_components") or ()) planned_shared_weights = component_attributes.get("workflow_planned_shared_weights") + planned_deferrals = component_attributes.get("workflow_planned_deferred_shared_weights") deferred = [] for shared_weight in component_attributes.get("shared_weights") or (): @@ -424,6 +425,8 @@ def _defer_shared_weight_aliases( continue if planned_shared_weights is not None and shared_weight["name"] not in planned_shared_weights: continue + if planned_deferrals is not None and shared_weight["name"] not in planned_deferrals: + continue alias = next( (endpoint for endpoint in shared_weight.get("aliases") or () if endpoint["component"] == component_name), None, @@ -642,9 +645,8 @@ def _copy_existing_quantization_config(quantization_config) -> dict | None: return copied -def _get_validated_mixed_precision_info(model: HfModelHandler) -> dict | None: +def _validate_mixed_precision_info(attributes: Mapping) -> dict | None: """Validate and copy the mixed-precision metadata consumed by quantization passes.""" - attributes = model.model_attributes or {} if "mixed_precision_info" not in attributes: return None @@ -678,6 +680,10 @@ def _get_validated_mixed_precision_info(model: HfModelHandler) -> dict | None: } +def _get_validated_mixed_precision_info(model: HfModelHandler) -> dict | None: + return _validate_mixed_precision_info(model.model_attributes or {}) + + def validate_moe_quantization_requirement( model: HfModelHandler, config: type[BasePassConfig], @@ -867,9 +873,9 @@ def prepare_model( MoE-capable consumer does not opt in. """ - existing_qcfg = _copy_existing_quantization_config( - getattr(model.get_hf_model_config(), "quantization_config", None) - ) + hf_config = model.get_hf_model_config() + existing_qcfg = _copy_existing_quantization_config(getattr(hf_config, "quantization_config", None)) + existing_deferred = deepcopy(getattr(hf_config, "olive_deferred_shared_weights", None) or []) if existing_qcfg is not None and existing_qcfg.get("quant_method", None) != OliveHfQuantizationMethod.OLIVE: raise ValueError("Model has an existing quantization configuration that is not compatible with this pass.") # Deliberately checked twice: prepare_model fails before loading, while @@ -974,17 +980,30 @@ def prepare_model( if auto_vision: fresh_qcfg.quantize_vision = bool(owned_vision_towers) - wrapper.olive_deferred_shared_weights = ( - _defer_shared_weight_aliases( + if existing_qcfg is None: + wrapper.olive_deferred_shared_weights = _defer_shared_weight_aliases( root_model, component_attributes, fresh_qcfg, lm_head_name=lm_head_name, embeds_name=embeds_name, ) - if existing_qcfg is None - else [] - ) + else: + wrapper.olive_deferred_shared_weights = existing_deferred + for request in existing_deferred: + alias = request["alias"]["parameter"].removesuffix(".weight") + if request["alias"]["component"] != component_name: + raise ValueError(f"Deferred shared weight {request['name']!r} belongs to another component.") + alias_module = get_attr(root_model, alias) + weight = getattr(alias_module, "weight", None) + if weight is None or isinstance(weight.data, QuantTensor): + raise ValueError(f"Deferred shared weight {request['name']!r} has no float alias {alias!r}.") + if alias == lm_head_name: + fresh_qcfg.lm_head = False + elif alias == embeds_name: + fresh_qcfg.embeds = False + else: + raise ValueError(f"Deferred shared weight {request['name']!r} has unknown alias {alias!r}.") fresh_skip_patterns = list(getattr(fresh_qcfg, "modules_to_not_convert", None) or []) component_embedding_name_set = set(component_embedding_names) @@ -1196,7 +1215,11 @@ def get_quant_config( """ validate_moe_quantization_requirement(model, config, existing_quantization_config) - mp_info = _get_validated_mixed_precision_info(model) + return _quant_config_from_pass(config, _get_validated_mixed_precision_info(model)) + + +def _quant_config_from_pass(config: type[BasePassConfig], mp_info: dict | None) -> OliveHfQuantizationConfig: + """Resolve pass options and optional mixed-precision metadata consistently.""" quant_config = { "bits": config.bits, "symmetric": config.sym, diff --git a/olive/workflows/run/builds.py b/olive/workflows/run/builds.py index 2e1cc0398..c0f35f87b 100644 --- a/olive/workflows/run/builds.py +++ b/olive/workflows/run/builds.py @@ -138,21 +138,35 @@ def _paths_overlap(first: Path, second: Path) -> bool: return first == second or first in second.parents or second in first.parents -def _build_quantizes_shared_table(build_config: dict, parameter: str) -> bool: - """Whether a component build requests a tied token table from KQuant/RTN.""" +def _planned_shared_qargs(build_config: dict, parameter: str, category: str, attributes: dict) -> dict | None: + """Resolve a scoped shared endpoint's effective KQuant/RTN layout.""" + from olive.hardware.accelerator import DEFAULT_CPU_ACCELERATOR + from olive.passes.pytorch.kquant import KQuant + from olive.passes.pytorch.quant_utils import _quant_config_from_pass, _validate_mixed_precision_info + from olive.passes.pytorch.rtn import Rtn + module_name = parameter.removesuffix(".weight") + mp_info = _validate_mixed_precision_info(attributes) + defaults = (mp_info or {}).get("default") or {} + for configs in build_config["passes"].values(): for pass_config in configs if isinstance(configs, list) else [configs]: parsed = pass_config.model_dump() if hasattr(pass_config, "model_dump") else pass_config - if parsed["type"].lower() not in {"kquant", "rtn"}: + pass_cls = {"kquant": KQuant, "rtn": Rtn}.get(parsed["type"].lower()) + if pass_cls is None: continue - options = parsed.get("config") or parsed - if options.get("embeds") is False: + options = parsed.get("config") or {key: value for key, value in parsed.items() if key != "type"} + pass_options = pass_cls.generate_config(DEFAULT_CPU_ACCELERATOR, options, disable_search=True) + quantization = _quant_config_from_pass(pass_options, mp_info) + enabled = defaults.get(category) + if enabled is None: + enabled = getattr(pass_options, category, False) + if enabled is False: continue - if match_skip(module_name, options.get("modules_to_not_convert") or []): + if match_skip(module_name, quantization.modules_to_not_convert or []): continue - return True - return False + return quantization.get_qlinear_init_args(module_name) + return None def expand_builds(run_config: dict) -> OrderedDict[str, dict]: @@ -217,17 +231,37 @@ def expand_builds(run_config: dict) -> OrderedDict[str, dict]: attributes = input_model["config"].get("model_attributes") or {} if attributes.get("shared_weights"): hf_builds.append((child_config, attributes)) - planned_shared_weights = set() + canonical_qargs = {} for child_config, attributes in hf_builds: for shared_weight in attributes["shared_weights"]: if shared_weight["kind"] != "tied_word_embeddings": continue if shared_weight["canonical"]["component"] != attributes.get("component_name"): continue - if _build_quantizes_shared_table(child_config, shared_weight["canonical"]["parameter"]): - planned_shared_weights.add(shared_weight["name"]) + qargs = _planned_shared_qargs(child_config, shared_weight["canonical"]["parameter"], "embeds", attributes) + if qargs is not None: + canonical_qargs[shared_weight["name"]] = qargs + deferred_shared_weights = set() + for child_config, attributes in hf_builds: + for shared_weight in attributes["shared_weights"]: + if shared_weight["name"] not in canonical_qargs: + continue + for alias in shared_weight["aliases"]: + if alias["component"] != attributes.get("component_name"): + continue + qargs = _planned_shared_qargs(child_config, alias["parameter"], "lm_head", attributes) + if qargs is None: + continue + if qargs != canonical_qargs[shared_weight["name"]]: + raise ValueError( + f"Shared weight {shared_weight['name']!r} has incompatible quantization layouts: " + f"{alias['parameter']}={qargs}, " + f"{shared_weight['canonical']['parameter']}={canonical_qargs[shared_weight['name']]}" + ) + deferred_shared_weights.add(shared_weight["name"]) for _, attributes in hf_builds: - attributes["workflow_planned_shared_weights"] = sorted(planned_shared_weights) + attributes["workflow_planned_shared_weights"] = sorted(canonical_qargs) + attributes["workflow_planned_deferred_shared_weights"] = sorted(deferred_shared_weights) return expanded diff --git a/olive/workflows/run/hf_component_assembly.py b/olive/workflows/run/hf_component_assembly.py index 1f3666357..61c0fe1b7 100644 --- a/olive/workflows/run/hf_component_assembly.py +++ b/olive/workflows/run/hf_component_assembly.py @@ -52,6 +52,7 @@ "codec_head", "embed_tokens", "lm_head", + "output", "output_projection", "proj_out", "shared", @@ -68,6 +69,7 @@ _OUTPUT_HEAD_MODULE_NAMES = { "codec_head", "lm_head", + "output", "output_projection", "proj_out", } @@ -265,7 +267,18 @@ def _resolve_shared_weights( } resolved = [] for shared_weight in declarations.values(): + if shared_weight.kind != "tied_word_embeddings": + raise ValueError(f"HF component assembly does not support shared weight kind {shared_weight.kind!r}.") endpoints = shared_weight.endpoints + for endpoint in endpoints: + artifact = artifacts_by_component.get(endpoint.component) + if artifact is None: + continue + if not _matches_source_path(_module_name(endpoint.parameter), artifact.source_paths): + raise ValueError( + f"Shared weight {shared_weight.name!r} endpoint {endpoint.parameter!r} " + f"is outside component {endpoint.component!r} source paths." + ) if any(endpoint.component not in artifacts_by_component for endpoint in endpoints): continue @@ -658,11 +671,52 @@ def _component_quantization_mapping( return mapping +def _float_shared_alias_sources(artifacts: list[_BuildArtifact]) -> dict[str, dict[str, str]]: + """Recover float aliases when only the canonical tied table was quantized.""" + artifacts_by_component = {component: artifact for artifact in artifacts for component in artifact.components} + sources: dict[str, dict[str, str]] = {} + for artifact in artifacts: + for shared_weight in artifact.shared_weights: + if shared_weight.kind != "tied_word_embeddings": + continue + canonical = shared_weight.canonical + canonical_artifact = artifacts_by_component.get(canonical.component) + if canonical_artifact is None or f"{canonical.parameter}_qweight" not in canonical_artifact.checkpoint.keys: + continue + for alias in shared_weight.aliases: + alias_artifact = artifacts_by_component.get(alias.component) + if alias_artifact is None or alias.parameter in alias_artifact.checkpoint.keys: + continue + if f"{alias.parameter}_qweight" in alias_artifact.checkpoint.keys: + continue + if any( + request.get("name") == shared_weight.name + for request in alias_artifact.config.get("olive_deferred_shared_weights", ()) + ): + continue + quantization = alias_artifact.config.get("quantization_config") or {} + if quantization.get("lm_head"): + raise ValueError( + f"Shared weight {shared_weight.name!r} requested quantization of " + f"{alias.parameter!r}, but build {alias_artifact.name!r} produced no packed alias." + ) + if canonical.parameter not in alias_artifact.checkpoint.keys: + raise ValueError( + f"Shared weight {shared_weight.name!r} is missing float source " + f"{canonical.parameter!r} for alias {alias.parameter!r} " + f"in build {alias_artifact.name!r}." + ) + sources.setdefault(alias_artifact.name, {})[alias.parameter] = canonical.parameter + return sources + + def _write_shards( entries: list[tuple[str, _Checkpoint]], output_dir: Path, prefix: str, relative_dir: Path, + *, + source_keys: dict[str, str] | None = None, ) -> tuple[dict[str, str], int]: output_dir.mkdir(parents=True, exist_ok=True) weight_map = {} @@ -684,7 +738,7 @@ def flush() -> None: batch_size = 0 for key, checkpoint in sorted(entries, key=lambda item: item[0]): - tensor = checkpoint.tensor(key) + tensor = checkpoint.tensor((source_keys or {}).get(key, key)) tensor_size = tensor.numel() * tensor.element_size() if batch and batch_size + tensor_size > _SHARD_LIMIT: flush() @@ -718,6 +772,7 @@ def _materialize_component_artifacts( owned_keys = set() exclusions = _shared_weight_exclusions(shared_weights or []) shared_weights = shared_weights or [] + float_alias_sources = _float_shared_alias_sources(artifacts) for artifact in artifacts: entries = [ @@ -725,6 +780,7 @@ def _materialize_component_artifacts( for key in artifact.checkpoint.keys if _matches_source_path(key, artifact.source_paths) and key not in exclusions.get(artifact.name, set()) ] + entries.extend((key, artifact.checkpoint) for key in float_alias_sources.get(artifact.name, {})) if not entries: raise ValueError( f"Build {artifact.name!r} source paths matched no checkpoint tensors: {artifact.source_paths}" @@ -741,6 +797,7 @@ def _materialize_component_artifacts( artifact_dir, "model", Path(artifact.name), + source_keys=float_alias_sources.get(artifact.name), ) weight_map.update(component_map) total_size += component_size @@ -999,6 +1056,7 @@ def try_assemble_hf_component_builds( "shared_weights", "workflow_components", "workflow_planned_shared_weights", + "workflow_planned_deferred_shared_weights", ): attributes.pop(name, None) attributes["assembled_components"] = [ diff --git a/olive/workflows/run/run.py b/olive/workflows/run/run.py index 6b6928daf..4ea5ba284 100644 --- a/olive/workflows/run/run.py +++ b/olive/workflows/run/run.py @@ -206,7 +206,16 @@ def _run_builds_in_parallel(package_config: OlivePackageConfig, parsed_config: M from olive.workflows.run.hf_component_assembly import try_assemble_hf_component_builds - try_assemble_hf_component_builds(build_configs, results, parsed_config.output_dir) + assembled = try_assemble_hf_component_builds(build_configs, results, parsed_config.output_dir) + if assembled is None and any( + (config.input_model.config.get("model_attributes") or {}).get("workflow_planned_deferred_shared_weights") + for config in build_configs.values() + if config.input_model.type.lower() == "hfmodel" + ): + raise RuntimeError( + "Deferred shared weights require automatic Hugging Face component assembly, " + "but the selected builds did not produce compatible HF component outputs." + ) return OrderedDict((build_name, results[build_name]) for build_name in build_configs) diff --git a/test/common/test_mobius_utils.py b/test/common/test_mobius_utils.py index 708195701..d69f2201d 100644 --- a/test/common/test_mobius_utils.py +++ b/test/common/test_mobius_utils.py @@ -90,6 +90,34 @@ def test_coerce_reads_cross_component_shared_weights(): ] +@pytest.mark.parametrize( + ("aliases", "message"), + [ + ([], "at least one alias"), + ([{"component": "decoder", "parameter": "model.embed_tokens.weight"}], "duplicate parameter"), + ( + [ + {"component": "decoder", "parameter": "lm_head.weight"}, + {"component": "decoder", "parameter": "lm_head.weight"}, + ], + "duplicate parameter", + ), + ], +) +def test_coerce_rejects_invalid_shared_weight_endpoints(aliases, message): + with pytest.raises(ValueError, match=message): + SharedWeightInfo.coerce( + { + "name": "word_embeddings", + "canonical": { + "component": "embedding", + "parameter": "model.embed_tokens.weight", + }, + "aliases": aliases, + } + ) + + def test_coerce_falls_back_to_legacy_kind_and_source_path(): # Older mobius releases expose ``kind``/``source_path`` (singular string). component = ComponentInfo.coerce( diff --git a/test/passes/pytorch/test_kquant.py b/test/passes/pytorch/test_kquant.py index 7f15464c1..41ab36e00 100644 --- a/test/passes/pytorch/test_kquant.py +++ b/test/passes/pytorch/test_kquant.py @@ -348,6 +348,42 @@ def test_kquant_defers_noncanonical_cross_component_tied_weight(tmp_path: Path): ] +def test_followup_rtn_preserves_deferred_lm_head(tmp_path: Path): + model_path = tmp_path / "input_model" + _make_local_tiny_tied_llama(model_path) + decoder_model = HfModelHandler( + model_path=str(model_path), + model_attributes={ + "component_name": "decoder", + "component_role": "decoder", + "component_source_paths": ["model.layers", "model.norm", "lm_head"], + "shared_weights": [tied_word_embedding_group()], + "workflow_components": ["decoder", "embedding"], + "workflow_planned_shared_weights": ["word_embeddings"], + }, + ) + first = create_pass_from_dict( + KQuant, + { + "bits": 4, + "group_size": 16, + "sym": True, + "overrides": {"lm_head": {"bits": 8}}, + }, + disable_search=True, + ).run(decoder_model, str(tmp_path / "first")) + first_deferred = first.get_hf_model_config().olive_deferred_shared_weights + + second = create_pass_from_dict(Rtn, {"bits": 4, "group_size": 16, "sym": True}, disable_search=True).run( + first, str(tmp_path / "second") + ) + reloaded = second.load_model() + + assert not isinstance(reloaded.lm_head.weight.data, QuantTensor) + assert reloaded.config.quantization_config.lm_head is False + assert reloaded.config.olive_deferred_shared_weights == first_deferred + + def test_kquant_quantizes_alias_when_canonical_component_is_not_built( tmp_path: Path, ): diff --git a/test/workflows/test_hf_component_assembly.py b/test/workflows/test_hf_component_assembly.py index c55f4f01d..37f1dc8fd 100644 --- a/test/workflows/test_hf_component_assembly.py +++ b/test/workflows/test_hf_component_assembly.py @@ -15,11 +15,13 @@ from safetensors.torch import save_file from transformers import AutoConfig +from olive.common.mobius_utils import SharedWeightInfo from olive.workflows.run import hf_component_assembly as assembly_module from olive.workflows.run.hf_component_assembly import ( _assembly_lock, _Checkpoint, _merge_quantization_config, + _resolve_shared_weights, _validate_build_compatibility, try_assemble_hf_component_builds, ) @@ -152,6 +154,38 @@ def metadata(self, key): return self._metadata[key] +@pytest.mark.parametrize( + ("kind", "source_paths", "message"), + [ + ("parameter_alias", ["lm_head"], "does not support shared weight kind"), + ("tied_word_embeddings", ["model.layers"], "outside component"), + ], +) +def test_shared_weight_assembly_rejects_unsupported_or_unowned_endpoints(kind, source_paths, message): + declaration = tied_word_embedding_group() + declaration["kind"] = kind + shared_weight = SharedWeightInfo.coerce(declaration) + decoder = SimpleNamespace( + name="decoder", + components=["decoder"], + source_paths=source_paths, + shared_weights=[shared_weight], + checkpoint=_MetadataCheckpoint({}), + config={}, + ) + embedding = SimpleNamespace( + name="embedding", + components=["embedding"], + source_paths=["model.embed_tokens"], + shared_weights=[shared_weight], + checkpoint=_MetadataCheckpoint({}), + config={}, + ) + + with pytest.raises(ValueError, match=message): + _resolve_shared_weights([decoder, embedding]) + + def _metadata_artifact(name, source_path, quantization_config, metadata, model_config=None): config = { "model_type": "llama", @@ -169,6 +203,52 @@ def _metadata_artifact(name, source_path, quantization_config, metadata, model_c ) +def test_shared_output_head_keeps_quantized_head_metadata(): + declaration = tied_word_embedding_group() + declaration["aliases"][0]["parameter"] = "output.weight" + shared_weight = SharedWeightInfo.coerce(declaration) + decoder = _metadata_artifact( + "decoder", + "output", + _quantization_config(group_size=16, symmetric=True, quantize_vision=False, skips=[]), + {}, + ) + embedding = _metadata_artifact( + "embedding", + "model.embed_tokens", + _quantization_config( + group_size=16, + symmetric=True, + quantize_vision=False, + skips=[], + embeds=True, + overrides={"model.embed_tokens": {"bits": 8}}, + ), + { + "model.embed_tokens.weight_qweight": ((8, 8), "U8"), + "model.embed_tokens.weight_scales": ((8, 1), "F32"), + }, + ) + resolved = [ + assembly_module._ResolvedSharedWeight( + info=shared_weight, + canonical_artifact=embedding, + alias_artifacts=[(shared_weight.aliases[0], decoder)], + sidecar_suffixes=("_qweight", "_scales"), + qargs={"bits": 8, "symmetric": True, "group_size": 16}, + ) + ] + + merged, _ = _merge_quantization_config([decoder, embedding], resolved) + component_configs = assembly_module._component_quantization_mapping([decoder, embedding], resolved) + + assert merged["lm_head"] is True + assert merged["embeds"] is True + assert merged["tie_word_embeddings"] is True + assert component_configs["decoder"]["lm_head"] is True + assert component_configs["decoder"]["overrides"]["output"] == {"bits": 8} + + def test_rejects_safetensors_shards_outside_checkpoint_root(tmp_path): checkpoint = tmp_path / "checkpoint" checkpoint.mkdir() @@ -307,6 +387,7 @@ def test_assembles_disjoint_hf_components_and_preserves_unbuilt_weights(tmp_path (None, None), ("deferred", None), ("missing_canonical", "produced no canonical packed tensor"), + ("missing_float_source", "missing float source"), ("layout", "incompatible quantization layouts"), ("tensor", "differs from canonical tensor"), ], @@ -328,8 +409,10 @@ def test_assembles_cross_component_tied_word_embeddings(tmp_path, failure, messa symmetric=True, quantize_vision=False, skips=[], - overrides=(None if failure in {"deferred", "missing_canonical"} else {"lm_head": {"bits": 8}}), - lm_head=failure not in {"deferred", "missing_canonical"}, + overrides=( + None if failure in {"deferred", "missing_canonical", "missing_float_source"} else {"lm_head": {"bits": 8}} + ), + lm_head=failure not in {"deferred", "missing_canonical", "missing_float_source"}, ) embedding_quantization = _quantization_config( group_size=32, @@ -373,7 +456,7 @@ def test_assembles_cross_component_tied_word_embeddings(tmp_path, failure, messa "model.language_model.layers.0.weight_scales": torch.ones(8, 2), "model.unoptimized.weight": torch.ones(8, 8), } - if failure not in {"deferred", "missing_canonical"}: + if failure not in {"deferred", "missing_canonical", "missing_float_source"}: decoder_tensors.update( { "lm_head.weight_qweight": tied_qweight, @@ -481,8 +564,11 @@ def test_assembles_cross_component_tied_word_embeddings(tmp_path, failure, messa assert embedding_manifest["quantization_config"]["tie_word_embeddings"] is True -@pytest.mark.parametrize("embedding_enabled", [True, False]) -def test_assembles_auto_selected_tied_component_quantization(tmp_path, embedding_enabled): +@pytest.mark.parametrize( + ("head_enabled", "embedding_enabled"), + [(True, True), (True, False), (False, True)], +) +def test_assembles_auto_selected_tied_component_quantization(tmp_path, head_enabled, embedding_enabled): from olive.common.quant.tensor import QuantTensor from olive.model import HfModelHandler from olive.passes.olive_pass import create_pass_from_dict @@ -491,6 +577,7 @@ def test_assembles_auto_selected_tied_component_quantization(tmp_path, embedding source = tmp_path / "source" make_local_tiny_dense_llama(source, tie_word_embeddings=True) + source_head = HfModelHandler(model_path=str(source)).load_model().lm_head.weight.detach().clone() shared_weights = [tied_word_embedding_group()] components = ["decoder", "embedding"] paths = { @@ -508,6 +595,8 @@ def test_assembles_auto_selected_tied_component_quantization(tmp_path, embedding } if not embedding_enabled: pass_configs["embedding"]["embeds"] = False + if not head_enabled: + pass_configs["decoder"]["lm_head"] = False planned_shared_weights = ["word_embeddings"] if embedding_enabled else [] build_configs = OrderedDict() results = OrderedDict() @@ -542,18 +631,21 @@ def test_assembles_auto_selected_tied_component_quantization(tmp_path, embedding assert try_assemble_hf_component_builds(build_configs, results, assembled_dir) == assembled_dir quantization = json.loads((assembled_dir / "config.json").read_text())["quantization_config"] - assert quantization["lm_head"] is True + assert quantization["lm_head"] is head_enabled assert quantization["embeds"] is embedding_enabled - assert quantization["tie_word_embeddings"] is embedding_enabled + assert quantization["tie_word_embeddings"] is (head_enabled and embedding_enabled) keys = _checkpoint_keys(assembled_dir) assert ("model.embed_tokens.weight_qweight" in keys) is embedding_enabled - assert ("lm_head.weight_qweight" in keys) is not embedding_enabled + assert ("lm_head.weight_qweight" in keys) is (head_enabled and not embedding_enabled) + assert ("lm_head.weight" in keys) is (not head_enabled and embedding_enabled) loaded = HfModelHandler(model_path=str(assembled_dir)).load_model() embedding = loaded.get_input_embeddings()._parameters["weight"] head = loaded.get_output_embeddings()._parameters["weight"] - assert (embedding is head) is embedding_enabled + assert (embedding is head) is (head_enabled and embedding_enabled) assert isinstance(embedding.data, QuantTensor) is embedding_enabled - assert isinstance(head.data, QuantTensor) + assert isinstance(head.data, QuantTensor) is head_enabled + if not head_enabled: + torch.testing.assert_close(head, source_head, rtol=0, atol=0) def test_quantization_merge_resolves_effective_overrides_and_float_skips(): diff --git a/test/workflows/test_run_builds.py b/test/workflows/test_run_builds.py index 9a638fac3..699e4624e 100644 --- a/test/workflows/test_run_builds.py +++ b/test/workflows/test_run_builds.py @@ -53,6 +53,28 @@ def _patch_engine_and_acc(self): accelerator_patch = patch.object(sys.modules[olive_run.__module__], "create_accelerator", return_value=acc_mock) return run_mock, acc_mock, engine_run_patch, accelerator_patch + @staticmethod + def _patch_tied_components(monkeypatch): + shared_weight = mobius_utils.SharedWeightInfo.coerce(tied_word_embedding_group()) + monkeypatch.setattr( + mobius_utils, + "inspect_components", + lambda *args, **kwargs: [ + mobius_utils.ComponentInfo( + name="decoder", + role="decoder", + source_paths=["model.layers", "model.norm", "lm_head"], + shared_weights=[shared_weight], + ), + mobius_utils.ComponentInfo( + name="embedding", + role="embedding", + source_paths=["model.embed_tokens"], + shared_weights=[shared_weight], + ), + ], + ) + def test_builds_tag_all_selected_hf_components(self, monkeypatch): monkeypatch.setattr( mobius_utils, @@ -93,40 +115,50 @@ def test_builds_tag_all_selected_hf_components(self, monkeypatch): assert attributes["workflow_components"] == ["decoder", "embedding"] @pytest.mark.parametrize( - ("embedding_options", "planned"), + ("embedding_options", "decoder_options", "mixed_default", "planned", "deferred", "error"), [ - ({}, ["word_embeddings"]), - ({"embeds": False}, []), - ({"modules_to_not_convert": ["model.embed_tokens"]}, []), + ({}, {}, None, ["word_embeddings"], ["word_embeddings"], None), + ({"embeds": False}, {}, None, [], [], None), + ({"modules_to_not_convert": ["model.embed_tokens"]}, {}, None, [], [], None), + ({}, {}, {"embeds": False}, [], [], None), + ({}, {"lm_head": False}, None, ["word_embeddings"], [], None), + ( + {"bits": 8}, + {"overrides": {"lm_head": {"bits": 8}}}, + None, + ["word_embeddings"], + ["word_embeddings"], + None, + ), + ({"bits": 8}, {}, None, None, None, "incompatible quantization layouts"), + ( + {"overrides": {"model.embed_tokens": {"bits": 8}}}, + {}, + None, + None, + None, + "incompatible quantization layouts", + ), ], ) - def test_builds_plan_only_quantized_shared_owners(self, monkeypatch, embedding_options, planned): - shared_weight = mobius_utils.SharedWeightInfo.coerce(tied_word_embedding_group()) - monkeypatch.setattr( - mobius_utils, - "inspect_components", - lambda *args, **kwargs: [ - mobius_utils.ComponentInfo( - name="decoder", - role="decoder", - source_paths=["model.layers", "lm_head"], - shared_weights=[shared_weight], - ), - mobius_utils.ComponentInfo( - name="embedding", - role="embedding", - source_paths=["model.embed_tokens"], - shared_weights=[shared_weight], - ), - ], - ) + def test_builds_plan_only_quantized_shared_owners( + self, monkeypatch, embedding_options, decoder_options, mixed_default, planned, deferred, error + ): + self._patch_tied_components(monkeypatch) config = deepcopy(self.template) config["input_model"] = { "type": "HfModel", - "config": {"model_path": "local/model"}, + "config": { + "model_path": "local/model", + **( + {"model_attributes": {"mixed_precision_info": {"default": mixed_default, "overrides": {}}}} + if mixed_default is not None + else {} + ), + }, } config["passes"] = { - "decoder_quant": {"type": "KQuant"}, + "decoder_quant": {"type": "KQuant", **decoder_options}, "embedding_quant": {"type": "KQuant", **embedding_options}, } config["builds"] = { @@ -140,42 +172,37 @@ def test_builds_plan_only_quantized_shared_owners(self, monkeypatch, embedding_o }, } + if error is not None: + with pytest.raises(ValueError, match=error): + expand_builds(config) + return + expanded = expand_builds(config) for build in expanded.values(): attributes = build["input_model"]["config"]["model_attributes"] assert attributes["workflow_components"] == ["decoder", "embedding"] assert attributes["workflow_planned_shared_weights"] == planned + assert attributes["workflow_planned_deferred_shared_weights"] == deferred - def test_builds_assemble_auto_selected_tied_weights(self, monkeypatch, tmp_path): + @pytest.mark.parametrize("mp_disables_embeds", [False, True]) + def test_builds_assemble_auto_selected_tied_weights(self, monkeypatch, tmp_path, mp_disables_embeds): from olive.common.quant.tensor import QuantTensor from olive.model import HfModelHandler from test.passes.pytorch.test_quantization_utils import make_local_tiny_dense_llama source = tmp_path / "source" make_local_tiny_dense_llama(source, tie_word_embeddings=True) - shared_weight = mobius_utils.SharedWeightInfo.coerce(tied_word_embedding_group()) - monkeypatch.setattr( - mobius_utils, - "inspect_components", - lambda *args, **kwargs: [ - mobius_utils.ComponentInfo( - name="decoder", - role="decoder", - source_paths=["model.layers", "model.norm", "lm_head"], - shared_weights=[shared_weight], - ), - mobius_utils.ComponentInfo( - name="embedding", - role="embedding", - source_paths=["model.embed_tokens"], - shared_weights=[shared_weight], - ), - ], - ) + self._patch_tied_components(monkeypatch) output_dir = tmp_path / "assembled" + model_attributes = ( + {"mixed_precision_info": {"default": {"embeds": False}, "overrides": {}}} if mp_disables_embeds else {} + ) config = { - "input_model": {"type": "HfModel", "config": {"model_path": str(source)}}, + "input_model": { + "type": "HfModel", + "config": {"model_path": str(source), "model_attributes": model_attributes}, + }, "passes": { "decoder_kquant": { "type": "KQuant", @@ -212,15 +239,40 @@ def test_builds_assemble_auto_selected_tied_weights(self, monkeypatch, tmp_path) olive_run(config) checkpoint = json.loads((output_dir / "model.safetensors.index.json").read_text(encoding="utf-8")) - assert "model.embed_tokens.weight_qweight" in checkpoint["weight_map"] - assert "lm_head.weight_qweight" not in checkpoint["weight_map"] + assert ("model.embed_tokens.weight_qweight" in checkpoint["weight_map"]) is not mp_disables_embeds + assert ("lm_head.weight_qweight" in checkpoint["weight_map"]) is mp_disables_embeds loaded = HfModelHandler(model_path=str(output_dir)).load_model() assert loaded.config.quantization_config.lm_head is True - assert loaded.config.quantization_config.embeds is True - assert loaded.config.quantization_config.tie_word_embeddings is True + assert loaded.config.quantization_config.embeds is not mp_disables_embeds + assert loaded.config.quantization_config.tie_word_embeddings is not mp_disables_embeds embedding = loaded.get_input_embeddings().weight - assert embedding is loaded.get_output_embeddings().weight - assert isinstance(embedding.data, QuantTensor) + assert (embedding is loaded.get_output_embeddings().weight) is not mp_disables_embeds + assert isinstance(embedding.data, QuantTensor) is not mp_disables_embeds + + def test_builds_reject_skipped_assembly_with_deferred_shared_weights(self, monkeypatch, tmp_path): + self._patch_tied_components(monkeypatch) + config = deepcopy(self.template) + config["input_model"] = {"type": "HfModel", "config": {"model_path": "local/model"}} + config["engine"]["output_dir"] = str(tmp_path / "assembled") + config["passes"] = { + "decoder_quant": {"type": "KQuant"}, + "embedding_quant": {"type": "KQuant"}, + "convert": {"type": "OnnxConversion"}, + } + config["builds"] = { + "decoder": {"components": ["decoder"], "pipeline": ["decoder_quant"]}, + "embedding": {"components": ["embedding"], "pipeline": ["embedding_quant"]}, + "onnx": {"pipeline": ["convert"]}, + } + output = MagicMock() + output.has_output_model.return_value = True + + with ( + patch.object(sys.modules[olive_run.__module__], "_run_single", return_value=output), + patch("olive.workflows.run.hf_component_assembly.try_assemble_hf_component_builds", return_value=None), + pytest.raises(RuntimeError, match="Deferred shared weights require automatic"), + ): + olive_run(config) def test_builds_components_on_non_composite_input_raises(self): config = deepcopy(self.template) From be1a942f9f31e79638668261784db9cbc5f5e70e Mon Sep 17 00:00:00 2001 From: Xiaoyu Zhang Date: Fri, 25 Sep 2026 16:05:33 -0700 Subject: [PATCH 9/9] Reload packed component embeddings and verify split export Build Hugging Face placeholders for component-owned secondary embeddings present in checkpoint sidecars, preventing their float reinitialization. Cover a tiny assembled tied Gemma4 checkpoint through Mobius embedding and decoder ONNX Runtime execution with full-logit parity. Signed-off-by: Xiaoyu Zhang --- .../configure-workflows/build-workflow.md | 1 + olive/common/quant/hf_utils.py | 9 ++ test/passes/pytorch/test_kquant.py | 38 +++++++ .../passes/pytorch/test_quantization_utils.py | 63 ++++++++++- test/workflows/test_tied_gemma4_export.py | 102 ++++++++++++++++++ 5 files changed, 209 insertions(+), 4 deletions(-) create mode 100644 test/workflows/test_tied_gemma4_export.py diff --git a/docs/source/how-to/configure-workflows/build-workflow.md b/docs/source/how-to/configure-workflows/build-workflow.md index 14dfe867c..b7c3afdba 100644 --- a/docs/source/how-to/configure-workflows/build-workflow.md +++ b/docs/source/how-to/configure-workflows/build-workflow.md @@ -187,6 +187,7 @@ already contain files. By default, each named build is saved under `/`. A build may set its own `output_dir` to any other location without changing where the assembled model is saved. Olive refuses to assemble into a workflow output directory that already contains files. +Tied embedding and LM-head builds must use matching quantization layouts; incompatible layouts fail before builds run. The named build directories contain component-only safetensors artifacts. The workflow output contains the complete checkpoint: diff --git a/olive/common/quant/hf_utils.py b/olive/common/quant/hf_utils.py index 0f45e3373..b766e3f1b 100644 --- a/olive/common/quant/hf_utils.py +++ b/olive/common/quant/hf_utils.py @@ -235,6 +235,14 @@ def _process_model_before_weight_loading( if keep_in_fp32_modules: skip_patterns.extend(keep_in_fp32_modules) + # Component builds may pack token tables beyond get_input_embeddings(). + extra_embeddings = ( + module + for name, module in model.named_modules() + if isinstance(module, nn.Embedding) + and self._checkpoint_keys is not None + and f"{name}.weight_qweight" in self._checkpoint_keys + ) for module, pname, full_name in iter_quant_targets( model, quantize_lm_head=self.quantization_config.lm_head, @@ -242,6 +250,7 @@ def _process_model_before_weight_loading( quantize_moe=self.quantization_config.moe, quantize_vision=getattr(self.quantization_config, "quantize_vision", False), skip_patterns=skip_patterns, + extra_embedding_modules=extra_embeddings, ): qargs = self.quantization_config.get_qlinear_init_args(full_name) param = module._parameters[pname] diff --git a/test/passes/pytorch/test_kquant.py b/test/passes/pytorch/test_kquant.py index 41ab36e00..cc9076570 100644 --- a/test/passes/pytorch/test_kquant.py +++ b/test/passes/pytorch/test_kquant.py @@ -7,6 +7,7 @@ import pytest import torch +from safetensors import safe_open from olive.common.quant.hf_utils import OliveHfQuantizationConfig from olive.common.quant.tensor import QuantTensor @@ -24,6 +25,7 @@ assert_dense_int2_mixed_precision_checkpoint, assert_uniform_int2_checkpoint, make_local_tiny_dense_llama, + make_local_tiny_tied_gemma4, plan_dense_int2_mixed_precision, tied_word_embedding_group, ) @@ -456,6 +458,42 @@ def test_component_embedding_selection_and_explicit_opt_out( assert not isinstance(loaded.lm_head._parameters["weight"].data, QuantTensor) +def test_component_embedding_reload_preserves_gemma4_per_layer_table(tmp_path: Path): + pytest.importorskip("transformers.models.gemma4") + source = tmp_path / "source" + make_local_tiny_tied_gemma4(source) + model = HfModelHandler( + model_path=str(source), + task="image-text-to-text", + model_attributes={ + "component_name": "embedding", + "component_role": "embedding", + "component_source_paths": [ + "model.language_model.embed_tokens", + "model.language_model.embed_tokens_per_layer", + "model.language_model.per_layer_model_projection", + "model.language_model.per_layer_projection_norm", + ], + }, + ) + output = tmp_path / "quantized" + quantizer = create_pass_from_dict(KQuant, {"bits": 8, "group_size": 16, "sym": True}, disable_search=True) + + loaded = quantizer.run(model, str(output)).load_model() + + table = loaded.model.language_model.embed_tokens_per_layer._parameters["weight"] + assert isinstance(table.data, QuantTensor) + assert not table.is_placeholder + assert loaded.config.quantization_config.embeds is True + with safe_open(output / "model.safetensors", framework="pt") as checkpoint: + torch.testing.assert_close( + table.qweight, + checkpoint.get_tensor("model.language_model.embed_tokens_per_layer.weight_qweight"), + rtol=0, + atol=0, + ) + + @pytest.mark.parametrize( "head_options", [ diff --git a/test/passes/pytorch/test_quantization_utils.py b/test/passes/pytorch/test_quantization_utils.py index 53ae64df3..6891c2b5f 100644 --- a/test/passes/pytorch/test_quantization_utils.py +++ b/test/passes/pytorch/test_quantization_utils.py @@ -34,8 +34,7 @@ def _dense_int2_calibration_dataset(seq_len: int, max_samples: int, vocab_size: def make_local_tiny_dense_llama(save_path: Path, *, tie_word_embeddings: bool = False) -> HfModelHandler: """Save a tiny dense Llama checkpoint and tokenizer without accessing the hub.""" - from tokenizers import Tokenizer, models, pre_tokenizers - from transformers import LlamaConfig, LlamaForCausalLM, PreTrainedTokenizerFast + from transformers import LlamaConfig, LlamaForCausalLM torch.manual_seed(0) save_path.mkdir(parents=True, exist_ok=True) @@ -49,11 +48,67 @@ def make_local_tiny_dense_llama(save_path: Path, *, tie_word_embeddings: bool = **({"tie_word_embeddings": True} if tie_word_embeddings else {}), ) LlamaForCausalLM(config).save_pretrained(save_path) + _save_trivial_tokenizer(save_path, config.vocab_size) + return HfModelHandler(model_path=str(save_path)) + - tokenizer = Tokenizer(models.WordLevel({f"t{i}": i for i in range(config.vocab_size)}, unk_token="t0")) +def make_local_tiny_tied_gemma4(save_path: Path) -> HfModelHandler: + """Save a tiny tied Gemma4 checkpoint with per-layer token embeddings.""" + from transformers import ( + Gemma4Config, + Gemma4ForConditionalGeneration, + Gemma4TextConfig, + Gemma4VisionConfig, + ) + + text = Gemma4TextConfig( # pylint: disable=unexpected-keyword-arg + vocab_size=128, + hidden_size=64, + intermediate_size=128, + num_hidden_layers=2, + num_attention_heads=4, + num_key_value_heads=2, + head_dim=16, + global_head_dim=16, + layer_types=["sliding_attention", "full_attention"], + hidden_size_per_layer_input=16, + vocab_size_per_layer_input=128, + num_kv_shared_layers=0, + attention_k_eq_v=False, + tie_word_embeddings=True, + ) + vision = Gemma4VisionConfig( # pylint: disable=unexpected-keyword-arg + hidden_size=32, + intermediate_size=64, + num_hidden_layers=1, + num_attention_heads=2, + num_key_value_heads=2, + head_dim=16, + patch_size=4, + position_embedding_size=64, + pooling_kernel_size=2, + ) + config = Gemma4Config( # pylint: disable=unexpected-keyword-arg + text_config=text, + vision_config=vision, + audio_config=None, + tie_word_embeddings=True, + image_token_id=120, + ) + torch.manual_seed(0) + save_path.mkdir(parents=True, exist_ok=True) + Gemma4ForConditionalGeneration(config).save_pretrained(save_path, save_original_format=False) + _save_trivial_tokenizer(save_path, config.text_config.vocab_size) + return HfModelHandler(model_path=str(save_path), task="image-text-to-text") + + +def _save_trivial_tokenizer(save_path: Path, vocab_size: int) -> None: + from tokenizers import Tokenizer, models, pre_tokenizers + from transformers import PreTrainedTokenizerFast + + tokenizer = Tokenizer(models.WordLevel({f"t{i}": i for i in range(vocab_size)}, unk_token="t0")) tokenizer.pre_tokenizer = pre_tokenizers.Whitespace() PreTrainedTokenizerFast(tokenizer_object=tokenizer, unk_token="t0", pad_token="t0").save_pretrained(save_path) - return HfModelHandler(model_path=str(save_path)) def tied_word_embedding_group(embedding_parameter: str = "model.embed_tokens.weight") -> dict: diff --git a/test/workflows/test_tied_gemma4_export.py b/test/workflows/test_tied_gemma4_export.py new file mode 100644 index 000000000..96664a304 --- /dev/null +++ b/test/workflows/test_tied_gemma4_export.py @@ -0,0 +1,102 @@ +# ------------------------------------------------------------------------- +# Copyright (c) Microsoft Corporation. All rights reserved. +# Licensed under the MIT License. +# -------------------------------------------------------------------------- +"""Network-free Olive -> Mobius split-export regression for tied Gemma4 weights.""" + +import json + +import numpy as np +import pytest +import torch + +from olive.model import HfModelHandler +from olive.workflows import run as olive_run +from test.passes.pytorch.test_quantization_utils import make_local_tiny_tied_gemma4 + + +def test_tied_gemma4_checkpoint_exports_and_runs_both_onnx_components(tmp_path): + pytest.importorskip("transformers.models.gemma4") + mobius = pytest.importorskip("mobius") + if not hasattr(mobius, "SharedWeightInfo"): + pytest.skip("Requires Mobius cross-component shared-weight support") + ir = pytest.importorskip("onnx_ir") + ort = pytest.importorskip("onnxruntime", minversion="1.28") + + source = tmp_path / "source" + make_local_tiny_tied_gemma4(source) + assembled = tmp_path / "assembled" + olive_run( + { + "input_model": { + "type": "HfModel", + "config": {"model_path": str(source), "task": "image-text-to-text"}, + }, + "passes": { + "decoder_kquant": { + "type": "KQuant", + "bits": 4, + "group_size": 16, + "sym": True, + "overrides": {"lm_head": {"bits": 8}}, + }, + "embedding_kquant": {"type": "KQuant", "bits": 8, "group_size": 16, "sym": True}, + }, + "builds": { + "decoder": {"components": ["decoder"], "pipeline": ["decoder_kquant"]}, + "embedding": {"components": ["embedding"], "pipeline": ["embedding_kquant"]}, + }, + "engine": { + "output_dir": str(assembled), + "cache_dir": str(tmp_path / "cache"), + "evaluate_input_model": False, + }, + "max_concurrent_builds": 1, + } + ) + + config = json.loads((assembled / "config.json").read_text(encoding="utf-8")) + index = json.loads((assembled / "model.safetensors.index.json").read_text(encoding="utf-8")) + assert config["quantization_config"]["tie_word_embeddings"] is True + assert "model.language_model.embed_tokens.weight_qweight" in index["weight_map"] + assert "lm_head.weight_qweight" not in index["weight_map"] + + package = mobius.build(str(assembled), dtype="f32") + sessions = {} + for component in ("embedding", "decoder"): + path = tmp_path / f"{component}.onnx" + ir.save(package[component], path, external_data=f"{component}.onnx.data") + sessions[component] = ort.InferenceSession(str(path), providers=["CPUExecutionProvider"]) + + ids = np.asarray([[2, 10]], dtype=np.int64) + embedding = sessions["embedding"] + embedding_outputs = dict( + zip( + (output.name for output in embedding.get_outputs()), + embedding.run(None, {"input_ids": ids, "image_features": np.zeros((0, 64), dtype=np.float32)}), + ) + ) + decoder = sessions["decoder"] + feeds = { + "inputs_embeds": embedding_outputs["inputs_embeds"], + "per_layer_inputs": embedding_outputs["per_layer_inputs"], + "attention_mask": np.ones((1, 2), dtype=np.int64), + "position_ids": np.asarray([[0, 1]], dtype=np.int64), + } + for input_info in decoder.get_inputs(): + if input_info.name.startswith("past_key_values."): + feeds[input_info.name] = np.zeros((1, 2, 0, 16), dtype=np.float32) + logits = decoder.run(["logits"], feeds)[0] + + hf_model = HfModelHandler(model_path=str(assembled), task="image-text-to-text").load_model().eval() + with torch.no_grad(): + reference = ( + hf_model( + input_ids=torch.from_numpy(ids), + attention_mask=torch.ones_like(torch.from_numpy(ids)), + use_cache=False, + ) + .logits.float() + .numpy() + ) + np.testing.assert_allclose(logits, reference, rtol=1e-3, atol=1e-3)