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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
17 changes: 8 additions & 9 deletions src/winml/modelkit/loader/config.py
Original file line number Diff line number Diff line change
Expand Up @@ -176,7 +176,7 @@ def resolve_loader_config(
"""
from transformers import AutoConfig

from .resolution import resolve_task
from .resolution import _architecture_loader_defaults, resolve_task

if trust_remote_code:
from ..utils.cli import warn_trust_remote_code
Expand Down Expand Up @@ -219,18 +219,14 @@ def resolve_loader_config(
f"attribute. Cannot proceed with config generation."
)

# Explicit model_type override alongside a model_id: thread the requested
# Explicit model_type: thread the requested
# variant through downstream resolution (task / class / composite tag and
# the loader config's model_type) WITHOUT mutating the loaded HF config. The
# exported graph, the htp Optimum patcher and every other consumer must keep
# seeing the architecture's native type; only the resolved build-variant tag
# changes. The model_type-only path above (AutoConfig.for_model) is
# unaffected because it only runs when model_id is None.
model_type_override = (
model_type
if (model_id is not None and model_type is not None and hf_config.model_type != model_type)
else None
)
# changes. Preserve even an explicit native type so checkpoint defaults
# cannot override it, including when hf_config was supplied by the caller.
model_type_override = model_type
if model_type_override is not None:
logger.info(
"Applying model_type override '%s' -> '%s' (explicit request)",
Expand All @@ -240,6 +236,7 @@ def resolve_loader_config(

# 2-3. Unified resolution. Task detection — including the no-architectures
# --model-type fallback (first supported task) — now lives in resolve_task.
architecture_defaults = _architecture_loader_defaults(hf_config, task, model_class, model_type)
resolution = resolve_task(
hf_config,
task=task,
Expand All @@ -260,6 +257,8 @@ def resolve_loader_config(
# resolved_hf_config keeps its native model_type.
if model_type_override is not None:
resolved_model_type = model_type_override
elif architecture_defaults is not None:
resolved_model_type = architecture_defaults[1]

# 5. Build loader config
loader_config = WinMLLoaderConfig(
Expand Down
20 changes: 14 additions & 6 deletions src/winml/modelkit/loader/hf.py
Original file line number Diff line number Diff line change
Expand Up @@ -236,11 +236,7 @@ def load_hf_model(
# freshly-loaded HF config. The torch model is instantiated from its own
# native config below, so export/patcher consumers keep the native type;
# only class/task resolution sees the variant.
model_type_override = (
model_type
if model_type is not None and getattr(hf_config, "model_type", None) != model_type
else None
)
model_type_override = model_type
if model_type_override is not None:
logger.info(
"Applying model_type override '%s' -> '%s' (explicit request)",
Expand Down Expand Up @@ -295,7 +291,19 @@ def load_hf_model(
load_kwargs["torch_dtype"] = torch_dtype
if attn_implementation is not None:
load_kwargs["attn_implementation"] = attn_implementation
model = loader_cls.from_pretrained(model_name_or_path, **load_kwargs)
if getattr(loader_cls, "_winml_require_complete_checkpoint", False) is True:
model, loading_info = loader_cls.from_pretrained(
model_name_or_path, output_loading_info=True, **load_kwargs
)
failures = {
key: loading_info[key]
for key in ("missing_keys", "unexpected_keys", "mismatched_keys", "error_msgs")
if loading_info.get(key)
}
if failures:
raise ValueError(f"Custom architecture has incompatible checkpoint weights: {failures}")
else:
model = loader_cls.from_pretrained(model_name_or_path, **load_kwargs)

# [5] Export Preparation
model.eval()
Expand Down
48 changes: 48 additions & 0 deletions src/winml/modelkit/loader/resolution.py
Original file line number Diff line number Diff line change
Expand Up @@ -250,6 +250,7 @@ class TaskSource(str, Enum):
USER_TASK = "user-task" # user passed --task
USER_CLASS = "user-class" # user passed --model-class; task inferred
MODEL_ID_DEFAULT = "model-id-default" # MODEL_TASK_MAPPING model-id default
ARCHITECTURE_DEFAULT = "architecture-default"
SENTINEL_DEFAULT = "sentinel-default" # (model_type, None) sentinel
TASKS_MANAGER = "tasks-manager" # Optimum inference (incl. fill-mask upgrade)
WRAPPED_LIBRARY = "wrapped-library" # no architectures -> first supported task
Expand Down Expand Up @@ -474,6 +475,40 @@ def _composite_display_class(model_type_norm: str, components: CompositeComponen
return cast("type", TasksManager.get_model_class_for_task(generation_first[0], framework="pt"))


def _architecture_loader_defaults(
config: PretrainedConfig,
task: str | None,
model_class: str | None,
model_type: str | None,
) -> tuple[str, str] | None:
"""Match declared architecture metadata without inspecting repository names."""
if model_class is not None or model_type is not None:
return None
from ..models.hf import ARCHITECTURE_LOADER_DEFAULTS

native_type = getattr(config, "model_type", None)
architectures = getattr(config, "architectures", None)
if not isinstance(native_type, str) or not isinstance(architectures, (list, tuple)):
return None
matches = [
ARCHITECTURE_LOADER_DEFAULTS[(native_type, name)]
for name in architectures
if isinstance(name, str) and (native_type, name) in ARCHITECTURE_LOADER_DEFAULTS
]
if not matches:
return None
if task is not None and all(normalize_task(task) != match[0] for match in matches):
return None
if len(architectures) != 1 or len(matches) != 1:
raise ValueError("Ambiguous registered architectures; specify loader overrides explicitly")
default_task, variant, requirements = matches[0]
if task is not None and normalize_task(task) != default_task:
return None
if any(getattr(config, key, None) != value for key, value in requirements.items()):
raise ValueError("Declared custom architecture has incompatible configuration")
return default_task, variant


def resolve_task(
config: PretrainedConfig,
*,
Expand All @@ -500,6 +535,19 @@ def resolve_task(

from optimum.exporters.tasks import TasksManager

defaults = _architecture_loader_defaults(config, task, model_class, model_type_override)
if defaults is not None:
default_task, variant = defaults
custom = _get_custom_model_class(variant.replace("_", "-"), default_task)
if custom is None:
raise ValueError(f"Checkpoint loader variant {variant!r} is not registered")
return TaskResolution(
default_task,
to_optimum_task(default_task),
custom,
TaskSource.USER_TASK if task is not None else TaskSource.ARCHITECTURE_DEFAULT,
)

model_type = model_type_override or getattr(config, "model_type", None)
model_type_norm = model_type.lower().replace("_", "-") if model_type else ""
model_id = getattr(config, "_name_or_path", "") or None
Expand Down
8 changes: 8 additions & 0 deletions src/winml/modelkit/models/hf/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -107,6 +107,7 @@
)
from .vision_encoder_decoder import VisionEncoderIOConfig as _VisionEncoderIOConfig
from .vitpose import MODEL_CLASS_MAPPING as _VITPOSE_CLASS_MAPPING
from .wav2vec2 import ARCHITECTURE_LOADER_DEFAULTS as _WAV2VEC2_LOADER_DEFAULTS
from .wav2vec2 import MODEL_CLASS_MAPPING as _WAV2VEC2_CLASS_MAPPING
from .wav2vec2 import (
# triggers registration
Expand Down Expand Up @@ -152,6 +153,12 @@
for _key, _model_cls in _sub_mapping.items()
}

# Declared architecture identities whose custom heads are absent from HF AutoModel.
# Values are (default task, registered build variant, required config fields), not tuned recipes.
ARCHITECTURE_LOADER_DEFAULTS: dict[tuple[str, str], tuple[str, str, dict[str, str]]] = {
**_WAV2VEC2_LOADER_DEFAULTS,
}

# Registry: model_type -> WinMLBuildConfig
# Only models that need non-autoconf-discoverable settings retain configs.
# Models with only optim flags rely on the analyzer autoconf loop.
Expand Down Expand Up @@ -183,6 +190,7 @@
}

__all__ = [
"ARCHITECTURE_LOADER_DEFAULTS",
"MODEL_BUILD_CONFIGS",
"MODEL_CLASS_MAPPING",
]
11 changes: 11 additions & 0 deletions src/winml/modelkit/models/hf/wav2vec2.py
Original file line number Diff line number Diff line change
Expand Up @@ -24,6 +24,15 @@

EMOTION_REGRESSION_MODEL_TYPE = "wav2vec2_emotion_regression"

# Architecture identity selects the custom head, never performance settings.
ARCHITECTURE_LOADER_DEFAULTS = {
("wav2vec2", "Wav2Vec2ForSpeechClassification"): (
"audio-classification",
EMOTION_REGRESSION_MODEL_TYPE,
{"problem_type": "regression"},
),
}


class RegressionHead(nn.Module):
"""Audeering dimensional-emotion regression head."""
Expand All @@ -46,6 +55,8 @@
class EmotionModel(Wav2Vec2PreTrainedModel):
"""Audeering wav2vec2 mean-pooling regression model."""

_winml_require_complete_checkpoint = True

def __init__(self, config: Any) -> None:
super().__init__(config)
self.config = config
Expand Down
Loading