diff --git a/src/winml/modelkit/loader/config.py b/src/winml/modelkit/loader/config.py index b2bb053b5..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 resolve_task + from .resolution import _architecture_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,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, @@ -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( diff --git a/src/winml/modelkit/loader/hf.py b/src/winml/modelkit/loader/hf.py index caf8dc615..f0f870ffc 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)", @@ -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() diff --git a/src/winml/modelkit/loader/resolution.py b/src/winml/modelkit/loader/resolution.py index 544d253a3..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,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, *, @@ -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 diff --git a/src/winml/modelkit/models/hf/__init__.py b/src/winml/modelkit/models/hf/__init__.py index 71fa3aee7..125e4befc 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 ARCHITECTURE_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() } +# 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. @@ -183,6 +190,7 @@ } __all__ = [ + "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 4e8259152..7c02dda32 100644 --- a/src/winml/modelkit/models/hf/wav2vec2.py +++ b/src/winml/modelkit/models/hf/wav2vec2.py @@ -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.""" @@ -46,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 new file mode 100644 index 000000000..f5bb192d3 --- /dev/null +++ b/tests/unit/loader/test_checkpoint_loader_defaults.py @@ -0,0 +1,199 @@ +# ------------------------------------------------------------------------- +# 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"], 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) + 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"], 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") + 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"], problem_type="regression" + ) + config._name_or_path = MODEL_ID + _, _, _, resolution = resolve_loader_config(MODEL_ID, hf_config=config) + assert resolution.source.value == "architecture-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" + + +@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"