From 8b409e1228362d06a20a78713173e97180068d22 Mon Sep 17 00:00:00 2001 From: Qiong Wu Date: Tue, 29 Sep 2026 17:36:20 +0800 Subject: [PATCH 1/2] fix(loader): resolve registered checkpoint heads without recipes --- src/winml/modelkit/loader/config.py | 20 ++-- src/winml/modelkit/loader/hf.py | 7 +- src/winml/modelkit/loader/resolution.py | 37 +++++++ src/winml/modelkit/models/hf/__init__.py | 8 ++ src/winml/modelkit/models/hf/wav2vec2.py | 8 ++ .../loader/test_checkpoint_loader_defaults.py | 104 ++++++++++++++++++ tests/unit/loader/test_load_hf_model.py | 12 +- 7 files changed, 179 insertions(+), 17 deletions(-) create mode 100644 tests/unit/loader/test_checkpoint_loader_defaults.py diff --git a/src/winml/modelkit/loader/config.py b/src/winml/modelkit/loader/config.py index b2bb053b5..35bd70d82 100644 --- a/src/winml/modelkit/loader/config.py +++ b/src/winml/modelkit/loader/config.py @@ -176,7 +176,7 @@ def resolve_loader_config( """ from transformers import AutoConfig - from .resolution import resolve_task + from .resolution import _checkpoint_loader_defaults, resolve_task if trust_remote_code: from ..utils.cli import warn_trust_remote_code @@ -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)", @@ -240,11 +236,15 @@ 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. + checkpoint_defaults = _checkpoint_loader_defaults( + hf_config, task, model_class, model_type, model_id + ) resolution = resolve_task( hf_config, task=task, model_class=model_class, model_type_override=model_type_override, + model_id=model_id, ) resolved_task, resolved_class = resolution.task, resolution.model_class logger.info("Resolved: task=%s, model_class=%s", resolved_task, resolved_class.__name__) @@ -260,6 +260,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 checkpoint_defaults is not None: + resolved_model_type = checkpoint_defaults[1] # 5. Build loader config loader_config = WinMLLoaderConfig( diff --git a/src/winml/modelkit/loader/hf.py b/src/winml/modelkit/loader/hf.py index caf8dc615..83e2beef4 100644 --- a/src/winml/modelkit/loader/hf.py +++ b/src/winml/modelkit/loader/hf.py @@ -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)", @@ -266,6 +262,7 @@ def load_hf_model( task=task, model_class=model_class, model_type_override=model_type_override, + model_id=model_name_or_path, ) task, resolved_class = resolution.task, resolution.model_class except ValueError as e: diff --git a/src/winml/modelkit/loader/resolution.py b/src/winml/modelkit/loader/resolution.py index 544d253a3..bb6dbcb4d 100644 --- a/src/winml/modelkit/loader/resolution.py +++ b/src/winml/modelkit/loader/resolution.py @@ -474,12 +474,34 @@ 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 _checkpoint_loader_defaults( + config: PretrainedConfig, + task: str | None, + model_class: str | None, + model_type: str | None, + model_id: str | None = None, +) -> tuple[str, str] | None: + """Resolve exact checkpoint identity without overriding explicit choices.""" + if model_class is not None or model_type is not None: + return None + from ..models.hf import CHECKPOINT_LOADER_DEFAULTS + + identity = model_id or getattr(config, "_name_or_path", "") + if not isinstance(identity, str): + return None + defaults = CHECKPOINT_LOADER_DEFAULTS.get(identity) + if defaults is None or (task is not None and normalize_task(task) != defaults[0]): + return None + return defaults + + def resolve_task( config: PretrainedConfig, *, task: str | None = None, model_class: str | None = None, model_type_override: str | None = None, + model_id: str | None = None, ) -> TaskResolution: """Resolve a single model's task + class from an HF config. @@ -490,6 +512,8 @@ def resolve_task( ``model_type_override`` lets a caller drive resolution with a build variant (e.g. ``qwen3_transformer_only``) without mutating the loaded HF config; when ``None`` the architecture's native ``config.model_type`` is used. + ``model_id`` supplies the requested checkpoint identity when a caller has + already loaded its config; otherwise ``config._name_or_path`` is used. """ if getattr(config, "_winml_generic_fallback", False) is True: raise ValueError( @@ -500,6 +524,19 @@ def resolve_task( from optimum.exporters.tasks import TasksManager + defaults = _checkpoint_loader_defaults(config, task, model_class, model_type_override, model_id) + 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.MODEL_ID_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 diff --git a/src/winml/modelkit/models/hf/__init__.py b/src/winml/modelkit/models/hf/__init__.py index 71fa3aee7..3651285b2 100644 --- a/src/winml/modelkit/models/hf/__init__.py +++ b/src/winml/modelkit/models/hf/__init__.py @@ -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 CHECKPOINT_LOADER_DEFAULTS as _WAV2VEC2_LOADER_DEFAULTS from .wav2vec2 import MODEL_CLASS_MAPPING as _WAV2VEC2_CLASS_MAPPING from .wav2vec2 import ( # triggers registration @@ -152,6 +153,12 @@ for _key, _model_cls in _sub_mapping.items() } +# Exact checkpoint identities whose custom heads are absent from HF AutoModel. +# Values are (default task, registered build variant), not tuned recipes. +CHECKPOINT_LOADER_DEFAULTS: dict[str, tuple[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. @@ -183,6 +190,7 @@ } __all__ = [ + "CHECKPOINT_LOADER_DEFAULTS", "MODEL_BUILD_CONFIGS", "MODEL_CLASS_MAPPING", ] diff --git a/src/winml/modelkit/models/hf/wav2vec2.py b/src/winml/modelkit/models/hf/wav2vec2.py index 4e8259152..acc8e63e4 100644 --- a/src/winml/modelkit/models/hf/wav2vec2.py +++ b/src/winml/modelkit/models/hf/wav2vec2.py @@ -24,6 +24,14 @@ EMOTION_REGRESSION_MODEL_TYPE = "wav2vec2_emotion_regression" +# Checkpoint identity selects the custom head, never performance settings. +CHECKPOINT_LOADER_DEFAULTS = { + "audeering/wav2vec2-large-robust-12-ft-emotion-msp-dim": ( + "audio-classification", + EMOTION_REGRESSION_MODEL_TYPE, + ), +} + class RegressionHead(nn.Module): """Audeering dimensional-emotion regression head.""" diff --git a/tests/unit/loader/test_checkpoint_loader_defaults.py b/tests/unit/loader/test_checkpoint_loader_defaults.py new file mode 100644 index 000000000..de1cef3db --- /dev/null +++ b/tests/unit/loader/test_checkpoint_loader_defaults.py @@ -0,0 +1,104 @@ +# ------------------------------------------------------------------------- +# Copyright (c) Microsoft Corporation. All rights reserved. +# Licensed under the MIT License. +# -------------------------------------------------------------------------- +"""Checkpoint loader identity must not depend on performance recipes.""" + +from unittest.mock import patch + +import pytest +from transformers import Wav2Vec2Config + +from winml.modelkit.loader import resolve_loader_config, resolve_task + + +MODEL_ID = "audeering/wav2vec2-large-robust-12-ft-emotion-msp-dim" + + +@pytest.mark.parametrize("task", [None, "audio-classification"]) +def test_checkpoint_loader_defaults_resolve_custom_head(task): + config = Wav2Vec2Config(architectures=["Wav2Vec2ForSpeechClassification"]) + config._name_or_path = MODEL_ID + with patch("winml.modelkit.loader._autoconfig.load_hf_config", return_value=config): + loader, resolved_config, cls, resolution = resolve_loader_config(MODEL_ID, task=task) + assert loader.task == "audio-classification" + assert loader.model_type == "wav2vec2_emotion_regression" + assert cls.__name__ == "EmotionModel" + assert resolution.model_class is cls + assert resolved_config.model_type == "wav2vec2" + direct = resolve_task(config, task=task) + assert direct.task == loader.task + assert direct.model_class is cls + + +def test_explicit_other_task_does_not_select_checkpoint_head(): + config = Wav2Vec2Config(architectures=["Wav2Vec2ForCTC"]) + config._name_or_path = MODEL_ID + result = resolve_task(config, task="automatic-speech-recognition") + assert result.model_class.__name__ != "EmotionModel" + + +def test_unregistered_checkpoint_does_not_select_custom_head(): + config = Wav2Vec2Config(architectures=["Wav2Vec2ForSequenceClassification"]) + config._name_or_path = "another/checkpoint" + result = resolve_task(config, task="audio-classification") + assert result.model_class.__name__ != "EmotionModel" + + +def test_explicit_model_type_keeps_native_loader(): + config = Wav2Vec2Config(architectures=["Wav2Vec2ForSequenceClassification"]) + config._name_or_path = MODEL_ID + loader, _, cls, _ = resolve_loader_config( + MODEL_ID, hf_config=config, model_type="wav2vec2", task="audio-classification" + ) + assert loader.model_type == "wav2vec2" + assert cls.__name__ != "EmotionModel" + + +def test_generated_build_config_keeps_default_optimizations(): + from winml.modelkit.config import generate_hf_build_config + + config = Wav2Vec2Config(architectures=["Wav2Vec2ForSpeechClassification"]) + config._name_or_path = MODEL_ID + with patch("winml.modelkit.loader._autoconfig.load_hf_config", return_value=config): + build = generate_hf_build_config(MODEL_ID, device="cpu", ep="cpu") + assert build.loader.model_class == "EmotionModel" + assert build.loader.task == "audio-classification" + assert build.loader.model_type == "wav2vec2_emotion_regression" + assert build.optim.to_dict() == {} + assert [tensor.name for tensor in build.export.input_tensors] == ["input_values"] + assert [tensor.name for tensor in build.export.output_tensors] == ["hidden_states", "logits"] + + +def test_preloaded_config_preserves_explicit_native_type(): + config = Wav2Vec2Config(architectures=["Wav2Vec2ForSequenceClassification"]) + config._name_or_path = MODEL_ID + loader, _, cls, _ = resolve_loader_config( + hf_config=config, model_type="wav2vec2", task="audio-classification" + ) + assert loader.model_type == "wav2vec2" + assert cls.__name__ != "EmotionModel" + + +def test_automatic_loader_provenance_is_not_user_task(): + config = Wav2Vec2Config(architectures=["Wav2Vec2ForSpeechClassification"]) + config._name_or_path = MODEL_ID + _, _, _, resolution = resolve_loader_config(MODEL_ID, hf_config=config) + assert resolution.source.value == "model-id-default" + + +def test_direct_load_preserves_explicit_model_type(): + from winml.modelkit.loader import load_hf_model + + config = Wav2Vec2Config(architectures=["Wav2Vec2ForSequenceClassification"]) + config._name_or_path = MODEL_ID + with ( + patch( + "winml.modelkit.loader.resolution.resolve_task", side_effect=RuntimeError("stop") + ) as resolve, + pytest.raises(RuntimeError, match="stop"), + ): + load_hf_model( + MODEL_ID, hf_config=config, model_type="wav2vec2", task="audio-classification" + ) + assert resolve.call_args.kwargs["model_type_override"] == "wav2vec2" diff --git a/tests/unit/loader/test_load_hf_model.py b/tests/unit/loader/test_load_hf_model.py index 01b46e22b..2f817cf60 100644 --- a/tests/unit/loader/test_load_hf_model.py +++ b/tests/unit/loader/test_load_hf_model.py @@ -116,7 +116,9 @@ def test_model_class_without_user_script_uses_tasks_manager(self, monkeypatch): # Track calls to resolve_task resolve_calls = [] - def mock_resolve(config, *, task=None, model_class=None, model_type_override=None): + def mock_resolve( + config, *, task=None, model_class=None, model_type_override=None, model_id=None + ): resolve_calls.append({"task": task, "model_class": model_class}) mock_class = MagicMock() mock_class.__name__ = "MockModel" @@ -162,7 +164,9 @@ def test_auto_detect_when_no_model_class(self, monkeypatch): # Track calls to resolve_task resolve_calls = [] - def mock_resolve(config, *, task=None, model_class=None, model_type_override=None): + def mock_resolve( + config, *, task=None, model_class=None, model_type_override=None, model_id=None + ): resolve_calls.append({"task": task, "model_class": model_class}) mock_class = MagicMock() mock_class.__name__ = "AutoDetectedModel" @@ -279,7 +283,9 @@ def test_bert_tiny_uses_model_specific_default_task(self, monkeypatch): resolve_calls = [] - def mock_resolve(config, *, task=None, model_class=None, model_type_override=None): + def mock_resolve( + config, *, task=None, model_class=None, model_type_override=None, model_id=None + ): resolved_task = task or "feature-extraction" resolve_calls.append({"task": resolved_task, "model_class": model_class}) mock_class = MagicMock() From 1c773f783b6b0e47e5c273012376fbb8ab3d2f5d Mon Sep 17 00:00:00 2001 From: Qiong Wu Date: Tue, 29 Sep 2026 17:51:31 +0800 Subject: [PATCH 2/2] fix(loader): resolve declared architectures and validate custom weights --- src/winml/modelkit/loader/config.py | 11 +- src/winml/modelkit/loader/hf.py | 15 ++- src/winml/modelkit/loader/resolution.py | 39 ++++--- src/winml/modelkit/models/hf/__init__.py | 10 +- src/winml/modelkit/models/hf/wav2vec2.py | 9 +- .../loader/test_checkpoint_loader_defaults.py | 103 +++++++++++++++++- tests/unit/loader/test_load_hf_model.py | 12 +- 7 files changed, 155 insertions(+), 44 deletions(-) diff --git a/src/winml/modelkit/loader/config.py b/src/winml/modelkit/loader/config.py index 35bd70d82..7a1c3b935 100644 --- a/src/winml/modelkit/loader/config.py +++ b/src/winml/modelkit/loader/config.py @@ -176,7 +176,7 @@ def resolve_loader_config( """ from transformers import AutoConfig - from .resolution import _checkpoint_loader_defaults, resolve_task + from .resolution import _architecture_loader_defaults, resolve_task if trust_remote_code: from ..utils.cli import warn_trust_remote_code @@ -236,15 +236,12 @@ 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. - checkpoint_defaults = _checkpoint_loader_defaults( - hf_config, task, model_class, model_type, model_id - ) + architecture_defaults = _architecture_loader_defaults(hf_config, task, model_class, model_type) resolution = resolve_task( hf_config, task=task, model_class=model_class, model_type_override=model_type_override, - model_id=model_id, ) resolved_task, resolved_class = resolution.task, resolution.model_class logger.info("Resolved: task=%s, model_class=%s", resolved_task, resolved_class.__name__) @@ -260,8 +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 checkpoint_defaults is not None: - resolved_model_type = checkpoint_defaults[1] + elif architecture_defaults is not None: + resolved_model_type = architecture_defaults[1] # 5. Build loader config loader_config = WinMLLoaderConfig( diff --git a/src/winml/modelkit/loader/hf.py b/src/winml/modelkit/loader/hf.py index 83e2beef4..f0f870ffc 100644 --- a/src/winml/modelkit/loader/hf.py +++ b/src/winml/modelkit/loader/hf.py @@ -262,7 +262,6 @@ def load_hf_model( task=task, model_class=model_class, model_type_override=model_type_override, - model_id=model_name_or_path, ) task, resolved_class = resolution.task, resolution.model_class except ValueError as e: @@ -292,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() diff --git a/src/winml/modelkit/loader/resolution.py b/src/winml/modelkit/loader/resolution.py index bb6dbcb4d..91e95994c 100644 --- a/src/winml/modelkit/loader/resolution.py +++ b/src/winml/modelkit/loader/resolution.py @@ -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 @@ -474,25 +475,38 @@ 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 _checkpoint_loader_defaults( +def _architecture_loader_defaults( config: PretrainedConfig, task: str | None, model_class: str | None, model_type: str | None, - model_id: str | None = None, ) -> tuple[str, str] | None: - """Resolve exact checkpoint identity without overriding explicit choices.""" + """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 CHECKPOINT_LOADER_DEFAULTS + from ..models.hf import ARCHITECTURE_LOADER_DEFAULTS - identity = model_id or getattr(config, "_name_or_path", "") - if not isinstance(identity, str): + 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 - defaults = CHECKPOINT_LOADER_DEFAULTS.get(identity) - if defaults is None or (task is not None and normalize_task(task) != defaults[0]): + 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 - return defaults + 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( @@ -501,7 +515,6 @@ def resolve_task( task: str | None = None, model_class: str | None = None, model_type_override: str | None = None, - model_id: str | None = None, ) -> TaskResolution: """Resolve a single model's task + class from an HF config. @@ -512,8 +525,6 @@ def resolve_task( ``model_type_override`` lets a caller drive resolution with a build variant (e.g. ``qwen3_transformer_only``) without mutating the loaded HF config; when ``None`` the architecture's native ``config.model_type`` is used. - ``model_id`` supplies the requested checkpoint identity when a caller has - already loaded its config; otherwise ``config._name_or_path`` is used. """ if getattr(config, "_winml_generic_fallback", False) is True: raise ValueError( @@ -524,7 +535,7 @@ def resolve_task( from optimum.exporters.tasks import TasksManager - defaults = _checkpoint_loader_defaults(config, task, model_class, model_type_override, model_id) + 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) @@ -534,7 +545,7 @@ def resolve_task( default_task, to_optimum_task(default_task), custom, - TaskSource.USER_TASK if task is not None else TaskSource.MODEL_ID_DEFAULT, + TaskSource.USER_TASK if task is not None else TaskSource.ARCHITECTURE_DEFAULT, ) model_type = model_type_override or getattr(config, "model_type", None) diff --git a/src/winml/modelkit/models/hf/__init__.py b/src/winml/modelkit/models/hf/__init__.py index 3651285b2..125e4befc 100644 --- a/src/winml/modelkit/models/hf/__init__.py +++ b/src/winml/modelkit/models/hf/__init__.py @@ -107,7 +107,7 @@ ) from .vision_encoder_decoder import VisionEncoderIOConfig as _VisionEncoderIOConfig from .vitpose import MODEL_CLASS_MAPPING as _VITPOSE_CLASS_MAPPING -from .wav2vec2 import CHECKPOINT_LOADER_DEFAULTS as _WAV2VEC2_LOADER_DEFAULTS +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 @@ -153,9 +153,9 @@ for _key, _model_cls in _sub_mapping.items() } -# Exact checkpoint identities whose custom heads are absent from HF AutoModel. -# Values are (default task, registered build variant), not tuned recipes. -CHECKPOINT_LOADER_DEFAULTS: dict[str, tuple[str, str]] = { +# 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, } @@ -190,7 +190,7 @@ } __all__ = [ - "CHECKPOINT_LOADER_DEFAULTS", + "ARCHITECTURE_LOADER_DEFAULTS", "MODEL_BUILD_CONFIGS", "MODEL_CLASS_MAPPING", ] diff --git a/src/winml/modelkit/models/hf/wav2vec2.py b/src/winml/modelkit/models/hf/wav2vec2.py index acc8e63e4..7c02dda32 100644 --- a/src/winml/modelkit/models/hf/wav2vec2.py +++ b/src/winml/modelkit/models/hf/wav2vec2.py @@ -24,11 +24,12 @@ EMOTION_REGRESSION_MODEL_TYPE = "wav2vec2_emotion_regression" -# Checkpoint identity selects the custom head, never performance settings. -CHECKPOINT_LOADER_DEFAULTS = { - "audeering/wav2vec2-large-robust-12-ft-emotion-msp-dim": ( +# Architecture identity selects the custom head, never performance settings. +ARCHITECTURE_LOADER_DEFAULTS = { + ("wav2vec2", "Wav2Vec2ForSpeechClassification"): ( "audio-classification", EMOTION_REGRESSION_MODEL_TYPE, + {"problem_type": "regression"}, ), } @@ -54,6 +55,8 @@ def forward(self, features: torch.Tensor) -> torch.Tensor: # noqa: D102 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 diff --git a/tests/unit/loader/test_checkpoint_loader_defaults.py b/tests/unit/loader/test_checkpoint_loader_defaults.py index de1cef3db..f5bb192d3 100644 --- a/tests/unit/loader/test_checkpoint_loader_defaults.py +++ b/tests/unit/loader/test_checkpoint_loader_defaults.py @@ -17,7 +17,9 @@ @pytest.mark.parametrize("task", [None, "audio-classification"]) def test_checkpoint_loader_defaults_resolve_custom_head(task): - config = Wav2Vec2Config(architectures=["Wav2Vec2ForSpeechClassification"]) + config = Wav2Vec2Config( + architectures=["Wav2Vec2ForSpeechClassification"], problem_type="regression" + ) config._name_or_path = MODEL_ID with patch("winml.modelkit.loader._autoconfig.load_hf_config", return_value=config): loader, resolved_config, cls, resolution = resolve_loader_config(MODEL_ID, task=task) @@ -58,7 +60,9 @@ def test_explicit_model_type_keeps_native_loader(): def test_generated_build_config_keeps_default_optimizations(): from winml.modelkit.config import generate_hf_build_config - config = Wav2Vec2Config(architectures=["Wav2Vec2ForSpeechClassification"]) + config = Wav2Vec2Config( + architectures=["Wav2Vec2ForSpeechClassification"], problem_type="regression" + ) config._name_or_path = MODEL_ID with patch("winml.modelkit.loader._autoconfig.load_hf_config", return_value=config): build = generate_hf_build_config(MODEL_ID, device="cpu", ep="cpu") @@ -81,10 +85,12 @@ def test_preloaded_config_preserves_explicit_native_type(): def test_automatic_loader_provenance_is_not_user_task(): - config = Wav2Vec2Config(architectures=["Wav2Vec2ForSpeechClassification"]) + config = Wav2Vec2Config( + architectures=["Wav2Vec2ForSpeechClassification"], problem_type="regression" + ) config._name_or_path = MODEL_ID _, _, _, resolution = resolve_loader_config(MODEL_ID, hf_config=config) - assert resolution.source.value == "model-id-default" + assert resolution.source.value == "architecture-default" def test_direct_load_preserves_explicit_model_type(): @@ -102,3 +108,92 @@ def test_direct_load_preserves_explicit_model_type(): MODEL_ID, hf_config=config, model_type="wav2vec2", task="audio-classification" ) assert resolve.call_args.kwargs["model_type_override"] == "wav2vec2" + + +@pytest.mark.parametrize("identity", ["other/renamed", "C:/models/local-copy"]) +def test_architecture_resolution_is_independent_of_model_id(identity): + config = Wav2Vec2Config( + architectures=["Wav2Vec2ForSpeechClassification"], problem_type="regression" + ) + config._name_or_path = identity + loader, _, cls, _ = resolve_loader_config(identity, hf_config=config) + assert cls.__name__ == "EmotionModel" + assert loader.model_type == "wav2vec2_emotion_regression" + + +def test_registered_architecture_rejects_incompatible_config(): + config = Wav2Vec2Config( + architectures=["Wav2Vec2ForSpeechClassification"], + problem_type="single_label_classification", + ) + with pytest.raises(ValueError, match="incompatible"): + resolve_loader_config("other/checkpoint", hf_config=config) + + +@pytest.mark.parametrize( + "field", ["missing_keys", "unexpected_keys", "mismatched_keys", "error_msgs"] +) +def test_custom_architecture_rejects_incompatible_weights(field): + from unittest.mock import MagicMock + + from winml.modelkit.loader import load_hf_model + from winml.modelkit.models.hf.wav2vec2 import EmotionModel + + config = Wav2Vec2Config( + architectures=["Wav2Vec2ForSpeechClassification"], problem_type="regression" + ) + info = { + name: [] for name in ["missing_keys", "unexpected_keys", "mismatched_keys", "error_msgs"] + } + info[field] = ["incompatible.weight"] + with ( + patch.object(EmotionModel, "from_pretrained", return_value=(MagicMock(), info)) as load, + pytest.raises(ValueError, match="incompatible checkpoint"), + ): + load_hf_model("renamed/model", hf_config=config) + assert load.call_args.kwargs["output_loading_info"] is True + + +def test_custom_architecture_accepts_complete_weight_load(): + from unittest.mock import MagicMock + + from winml.modelkit.loader import load_hf_model + from winml.modelkit.models.hf.wav2vec2 import EmotionModel + + config = Wav2Vec2Config( + architectures=["Wav2Vec2ForSpeechClassification"], problem_type="regression" + ) + model = MagicMock() + info = { + name: [] for name in ["missing_keys", "unexpected_keys", "mismatched_keys", "error_msgs"] + } + with patch.object(EmotionModel, "from_pretrained", return_value=(model, info)): + actual, _, task = load_hf_model("local-copy", hf_config=config) + assert actual is model + assert task == "audio-classification" + + +def test_explicit_task_bypasses_ambiguous_architecture_default(): + config = Wav2Vec2Config( + architectures=["Wav2Vec2ForSpeechClassification", "Wav2Vec2ForCTC"], + problem_type="regression", + ) + result = resolve_task(config, task="automatic-speech-recognition") + assert result.model_class.__name__ != "EmotionModel" + + +def test_ambiguous_architecture_requires_explicit_choice(): + config = Wav2Vec2Config( + architectures=["Wav2Vec2ForSpeechClassification", "Wav2Vec2ForCTC"], + problem_type="regression", + ) + with pytest.raises(ValueError, match="Ambiguous"): + resolve_task(config) + + +def test_explicit_native_type_bypasses_registered_architecture(): + config = Wav2Vec2Config( + architectures=["Wav2Vec2ForSpeechClassification"], problem_type="regression" + ) + result = resolve_task(config, task="audio-classification", model_type_override="wav2vec2") + assert result.model_class.__name__ != "EmotionModel" diff --git a/tests/unit/loader/test_load_hf_model.py b/tests/unit/loader/test_load_hf_model.py index 2f817cf60..01b46e22b 100644 --- a/tests/unit/loader/test_load_hf_model.py +++ b/tests/unit/loader/test_load_hf_model.py @@ -116,9 +116,7 @@ def test_model_class_without_user_script_uses_tasks_manager(self, monkeypatch): # Track calls to resolve_task resolve_calls = [] - def mock_resolve( - config, *, task=None, model_class=None, model_type_override=None, model_id=None - ): + def mock_resolve(config, *, task=None, model_class=None, model_type_override=None): resolve_calls.append({"task": task, "model_class": model_class}) mock_class = MagicMock() mock_class.__name__ = "MockModel" @@ -164,9 +162,7 @@ def test_auto_detect_when_no_model_class(self, monkeypatch): # Track calls to resolve_task resolve_calls = [] - def mock_resolve( - config, *, task=None, model_class=None, model_type_override=None, model_id=None - ): + def mock_resolve(config, *, task=None, model_class=None, model_type_override=None): resolve_calls.append({"task": task, "model_class": model_class}) mock_class = MagicMock() mock_class.__name__ = "AutoDetectedModel" @@ -283,9 +279,7 @@ def test_bert_tiny_uses_model_specific_default_task(self, monkeypatch): resolve_calls = [] - def mock_resolve( - config, *, task=None, model_class=None, model_type_override=None, model_id=None - ): + def mock_resolve(config, *, task=None, model_class=None, model_type_override=None): resolved_task = task or "feature-extraction" resolve_calls.append({"task": resolved_task, "model_class": model_class}) mock_class = MagicMock()