diff --git a/docs/source/how-to/configure-workflows/build-workflow.md b/docs/source/how-to/configure-workflows/build-workflow.md index d180cd108..b7c3afdba 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"] @@ -178,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: @@ -189,6 +199,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 ``` 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 8c4e5c13b..ef0e35d84 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 logger = logging.getLogger(__name__) @@ -28,6 +28,88 @@ 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 + + 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=data["component"], parameter=data["parameter"]) + duck_data: Any = data + return cls( + component=duck_data.component, + parameter=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" + + 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=data["name"], + canonical=SharedWeightEndpoint.coerce(data["canonical"]), + aliases=[SharedWeightEndpoint.coerce(alias) for alias in data.get("aliases", ())], + kind=data.get("kind", "parameter_alias"), + ) + duck_data: Any = data + return cls( + name=duck_data.name, + canonical=SharedWeightEndpoint.coerce(duck_data.canonical), + aliases=[SharedWeightEndpoint.coerce(alias) for alias in duck_data.aliases], + kind=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. @@ -41,6 +123,8 @@ 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. """ @@ -48,6 +132,7 @@ class ComponentInfo: role: Optional[str] = None source_paths: list[str] = 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": @@ -65,20 +150,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: 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/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/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/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 9d7111146..689f975f4 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,100 @@ 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, + 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 [] + 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 (): + if shared_weight.get("kind") != "tied_word_embeddings": + continue + canonical = shared_weight["canonical"] + if canonical["component"] == component_name: + 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 + 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, + ) + if alias is None: + 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" + 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, @@ -531,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 @@ -567,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], @@ -756,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 @@ -844,6 +961,50 @@ 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) + + 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, + ) + 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) extra_embedding_modules = component_embedding_modules.values() if component_source_paths else () @@ -870,6 +1031,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): @@ -941,6 +1108,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( @@ -1037,15 +1215,19 @@ 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, "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 {}, } @@ -1518,6 +1700,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/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 8325f8324..c0f35f87b 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,37 @@ def _paths_overlap(first: Path, second: Path) -> bool: return first == second or first in second.parents or second in first.parents +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 + pass_cls = {"kquant": KQuant, "rtn": Rtn}.get(parsed["type"].lower()) + if pass_cls is None: + continue + 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, quantization.modules_to_not_convert or []): + continue + return quantization.get_qlinear_init_args(module_name) + return None + + def expand_builds(run_config: dict) -> OrderedDict[str, dict]: """Expand ``builds`` into independent, ordinary Olive run configurations.""" if not isinstance(run_config, dict): @@ -154,6 +186,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,12 +212,56 @@ 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 + hf_builds = [] + 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"): + hf_builds.append((child_config, attributes)) + 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 + 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(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 40bca1e7d..61c0fe1b7 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"} @@ -50,12 +52,27 @@ "codec_head", "embed_tokens", "lm_head", + "output", "output_projection", "proj_out", "shared", "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", + "output_projection", + "proj_out", +} @dataclass @@ -70,6 +87,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: @@ -117,10 +144,36 @@ 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: + """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.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]): + 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 +230,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 +241,150 @@ 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(): + 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 + + 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 = [] + 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 +465,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 +492,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 +506,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 +538,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 +571,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,23 +624,99 @@ 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 +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 = {} @@ -427,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() @@ -453,18 +764,23 @@ 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 [] + float_alias_sources = _float_shared_alias_sources(artifacts) 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()) ] + 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}" @@ -481,16 +797,32 @@ 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 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 +1002,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 +1011,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 +1037,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 +1048,16 @@ 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", + "workflow_components", + "workflow_planned_shared_weights", + "workflow_planned_deferred_shared_weights", + ): attributes.pop(name, None) attributes["assembled_components"] = [ component for artifact in artifacts for component in artifact.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 872832be9..d69f2201d 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(): @@ -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( @@ -36,6 +45,79 @@ 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", + } + ], + } + ) + ] + + +@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/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..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 @@ -18,12 +19,15 @@ 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, assert_uniform_int2_checkpoint, make_local_tiny_dense_llama, + make_local_tiny_tied_gemma4, plan_dense_int2_mixed_precision, + tied_word_embedding_group, ) from test.utils import get_tiny_phi3 @@ -63,6 +67,11 @@ 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.""" + make_local_tiny_dense_llama(save_path, tie_word_embeddings=True) + + def _is_quant(module: torch.nn.Module) -> bool: if not isinstance(module, (torch.nn.Linear, torch.nn.Embedding)): return False @@ -259,6 +268,292 @@ 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 = [tied_word_embedding_group()] + 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, + "workflow_components": ["decoder", "embedding"], + }, + ) + 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, + "workflow_components": ["decoder", "embedding"], + }, + ) + decoder_pass = create_pass_from_dict( + KQuant, + { + "bits": 4, + "group_size": 16, + "sym": True, + "overrides": {"lm_head": {"bits": 8}}, + }, + disable_search=True, + ) + embedding_pass = create_pass_from_dict( + KQuant, + { + "bits": 8, + "group_size": 16, + "sym": 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 embedding.config.quantization_config.embeds is True + 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, + }, + } + ] + + +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, +): + 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": [tied_word_embedding_group()], + }, + ) + decoder_pass = create_pass_from_dict( + KQuant, + { + "bits": 4, + "group_size": 16, + "sym": 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 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) + + +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", + [ + {"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..56aa5ce70 100644 --- a/test/passes/pytorch/test_quant_utils.py +++ b/test/passes/pytorch/test_quant_utils.py @@ -1070,6 +1070,34 @@ 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() + 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, + 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(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..6891c2b5f 100644 --- a/test/passes/pytorch/test_quantization_utils.py +++ b/test/passes/pytorch/test_quantization_utils.py @@ -32,10 +32,9 @@ 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 + from transformers import LlamaConfig, LlamaForCausalLM torch.manual_seed(0) save_path.mkdir(parents=True, exist_ok=True) @@ -46,13 +45,83 @@ 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) + _save_trivial_tokenizer(save_path, config.vocab_size) + return HfModelHandler(model_path=str(save_path)) + + +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") + - tokenizer = Tokenizer(models.WordLevel({f"t{i}": i for i in range(config.vocab_size)}, unk_token="t0")) +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: + """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: diff --git a/test/workflows/test_hf_component_assembly.py b/test/workflows/test_hf_component_assembly.py index 70a8c924d..37f1dc8fd 100644 --- a/test/workflows/test_hf_component_assembly.py +++ b/test/workflows/test_hf_component_assembly.py @@ -15,14 +15,17 @@ 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, ) +from test.passes.pytorch.test_quantization_utils import tied_word_embedding_group def _quantization_config( @@ -69,15 +72,32 @@ 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, + role: str | None = None, + workflow_components: list[str] | 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 + 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", config={ - "model_attributes": { - "component_name": component, - "component_source_paths": [source_path], - } + "model_attributes": model_attributes, }, ), engine=SimpleNamespace(output_dir=output_dir), @@ -134,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", @@ -151,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() @@ -283,6 +381,273 @@ 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), + ("missing_canonical", "produced no canonical packed tensor"), + ("missing_float_source", "missing float source"), + ("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 = [tied_word_embedding_group("model.language_model.embed_tokens.weight")] + decoder_quantization = _quantization_config( + group_size=32, + symmetric=True, + quantize_vision=False, + skips=[], + 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, + symmetric=True, + quantize_vision=False, + skips=[], + overrides={ + "model.language_model.embed_tokens": { + "bits": 8, + **({"group_size": 16} if failure == "layout" else {}), + } + }, + embeds=failure != "missing_canonical", + ) + 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 in {"deferred", "missing_canonical"}: + 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 not in {"deferred", "missing_canonical", "missing_float_source"}: + 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, + ) + 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, + embedding_tensors, + 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 + + +@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 + 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) + source_head = HfModelHandler(model_path=str(source)).load_model().lm_head.weight.detach().clone() + shared_weights = [tied_word_embedding_group()] + 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 + if not head_enabled: + pass_configs["decoder"]["lm_head"] = 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 head_enabled + assert quantization["embeds"] 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 (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 (head_enabled and embedding_enabled) + assert isinstance(embedding.data, QuantTensor) is embedding_enabled + 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(): decoder_config = _quantization_config( group_size=32, diff --git a/test/workflows/test_run_builds.py b/test/workflows/test_run_builds.py index f38e100a3..699e4624e 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 @@ -12,7 +13,10 @@ 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.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 @@ -49,6 +53,227 @@ 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, + "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"] + + @pytest.mark.parametrize( + ("embedding_options", "decoder_options", "mixed_default", "planned", "deferred", "error"), + [ + ({}, {}, 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, 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", + **( + {"model_attributes": {"mixed_precision_info": {"default": mixed_default, "overrides": {}}}} + if mixed_default is not None + else {} + ), + }, + } + config["passes"] = { + "decoder_quant": {"type": "KQuant", **decoder_options}, + "embedding_quant": {"type": "KQuant", **embedding_options}, + } + config["builds"] = { + "decoder": { + "components": ["decoder"], + "pipeline": ["decoder_quant"], + }, + "embedding": { + "components": ["embedding"], + "pipeline": ["embedding_quant"], + }, + } + + 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 + + @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) + 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), "model_attributes": model_attributes}, + }, + "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"]) 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 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) 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) config["builds"] = { 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"] = { 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)