Skip to content
16 changes: 14 additions & 2 deletions docs/source/how-to/configure-workflows/build-workflow.md
Original file line number Diff line number Diff line change
Expand Up @@ -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": {
Expand All @@ -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"]
Expand All @@ -178,6 +187,7 @@ already contain files.
By default, each named build is saved under `<engine.output_dir>/<build-name>`. 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:
Expand All @@ -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
```
Expand Down
1 change: 1 addition & 0 deletions olive/common/hf/wrapper.py
Original file line number Diff line number Diff line change
Expand Up @@ -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]:
Expand Down
109 changes: 105 additions & 4 deletions olive/common/mobius_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,7 +14,7 @@

import logging
from dataclasses import dataclass, field
from typing import Optional
from typing import Any, Optional

logger = logging.getLogger(__name__)

Expand All @@ -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.
Expand All @@ -41,13 +123,16 @@ class ComponentInfo:
source_paths: Dotted submodule paths locating the component inside the full model
(e.g. ``["model.language_model"]``). A component may span multiple disjoint
sub-modules, so this is a list.
metadata: Additional component metadata retained from earlier callers.
shared_weights: Cross-component shared-weight declarations from Mobius.

"""

name: str
role: Optional[str] = None
source_paths: list[str] = field(default_factory=list)
metadata: dict = field(default_factory=dict)
Comment thread
Copilot marked this conversation as resolved.
shared_weights: list[SharedWeightInfo] = field(default_factory=list)

@classmethod
def coerce(cls, data: "ComponentInfo | dict | object") -> "ComponentInfo":
Expand All @@ -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 {}),
)


Expand Down
9 changes: 9 additions & 0 deletions olive/common/quant/hf_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -235,13 +235,22 @@ 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,
quantize_embeds=self.quantization_config.embeds,
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]
Expand Down
12 changes: 12 additions & 0 deletions olive/model/config/model_config.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)

Expand Down
2 changes: 1 addition & 1 deletion olive/passes/pytorch/kquant.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down
Loading
Loading