diff --git a/docs/reference/index.md b/docs/reference/index.md index 579abff2a..30541b579 100644 --- a/docs/reference/index.md +++ b/docs/reference/index.md @@ -55,6 +55,7 @@ stages based on the target device and precision. | `do_constant_folding` | `bool` | `true` | Fold constants during export. | | `verbose` | `bool` | `false` | Verbose export logging. | | `dynamo` | `bool` | `false` | Use PyTorch's TorchDynamo ONNX exporter; the default `false` selects the legacy TorchScript exporter. | +| `compatibility` | `dict \| null` | omitted | Resolved export compatibility knobs. Currently supports `transformers_attention: "eager"` to request eager Transformers attention during ONNX export. | | `enable_hierarchy_tags` | `bool` | `true` | Add module hierarchy tags to ONNX nodes. | | `clean_onnx` | `bool` | `false` | Strip hierarchy tags after export. | | `hierarchy_tag_format` | `"full" \| "module_only"` | `"full"` | Tag detail level. | diff --git a/pyproject.toml b/pyproject.toml index 4d3eca70a..dc860f182 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -199,6 +199,7 @@ include = [ "winml", "winml.*" ] "runtime_checker/need_rerun_errors.json", ] "winml.modelkit.export" = [ + "compatibility_rules.json", "htp/htp_metadata_schema.json", ] "winml.modelkit.pattern" = [ diff --git a/src/winml/modelkit/commands/build.py b/src/winml/modelkit/commands/build.py index b679d6095..733063492 100644 --- a/src/winml/modelkit/commands/build.py +++ b/src/winml/modelkit/commands/build.py @@ -901,7 +901,6 @@ def build( "EPNameOrAlias | None", _reject_ep_source(ep, "winml build"), ) - # Validate mutual exclusion if output_dir and use_cache: raise click.UsageError("--output-dir and --use-cache are mutually exclusive.") @@ -933,20 +932,6 @@ def build( dynamic_axes=dynamic_axes, ) - # Resolve an omitted EP and the requested device as one target. Forwarding - # both concrete axes keeps analyzer/build output aligned with the target - # selected by the catalog-backed resolver. - if ep_value is None: - from ..session import EPDeviceTarget, resolve_device - - try: - resolved_target = resolve_device(EPDeviceTarget(ep="auto", device=device)) - except ValueError as e: - raise click.UsageError(str(e)) from e - device = resolved_target.device - ep_value = cast("EPNameOrAlias", resolved_target.ep) - logger.info("Auto-resolved device=%s, EP=%s", device, ep_value) - try: # Hub-hosted ONNX (e.g. ``onnx-community/sam3-tracker-ONNX/onnx/...``) # is downloaded once and treated as a local .onnx file thereafter. @@ -958,6 +943,21 @@ def build( if model_input.kind is ModelInputKind.INVALID: raise click.UsageError(model_input.error or f"Invalid model input: {model}") + request_device = device + request_ep_value = ep_value + runtime_device = request_device + runtime_ep_value = request_ep_value + if runtime_ep_value is None: + from ..session import EPDeviceTarget, resolve_device + + try: + resolved_target = resolve_device(EPDeviceTarget(ep="auto", device=runtime_device)) + except ValueError as e: + raise click.UsageError(str(e)) from e + runtime_device = resolved_target.device + runtime_ep_value = cast("EPNameOrAlias", resolved_target.ep) + logger.info("Auto-resolved device=%s, EP=%s", runtime_device, runtime_ep_value) + # Load or auto-generate config if config_file is not None: config_or_configs = _load_config( @@ -989,6 +989,9 @@ def build( ] else: config_or_configs = merge_export_overrides(config_or_configs, export_overrides) + from ..config.build import apply_export_compatibility_policy + + apply_export_compatibility_policy(config_or_configs, device=device, ep=ep_value) else: if not model: raise click.UsageError("-m/--model is required when -c is not provided.") @@ -1013,17 +1016,18 @@ def build( ) config_or_configs = generate_build_config( onnx_path=model, - device=device, + device=runtime_device, precision=precision, - ep=ep_value, + ep=runtime_ep_value, ) else: config_or_configs = generate_build_config( model, trust_remote_code=trust_remote_code, - device=device, + device=runtime_device, precision=precision, - ep=ep_value, + ep=runtime_ep_value, + export_policy_target=(request_device, request_ep_value), shape_config=shape_overrides, override={"export": export_overrides} if export_overrides else None, ) @@ -1053,7 +1057,7 @@ def _patch_device(cfg: WinMLBuildConfig) -> None: from ..config import resolve_quant_compile_config resolved_quant, _ = resolve_quant_compile_config( - device=device, precision=precision, ep=ep_value + device=runtime_device, precision=precision, ep=runtime_ep_value ) if not quant or resolved_quant is None or is_pre_quantized_onnx_input: cfg.quant = None @@ -1081,7 +1085,7 @@ def _patch_device(cfg: WinMLBuildConfig) -> None: cfg.precision = precision.lower() # type: ignore[attr-defined] if cfg.compile is not None and cfg.compile.ep_config is not None: provider = cfg.compile.ep_config.provider - patched = WinMLCompileConfig.for_provider(provider, device=device) + patched = WinMLCompileConfig.for_provider(provider, device=runtime_device) if patched is not None: cfg.compile = patched @@ -1150,8 +1154,8 @@ def _patch_device(cfg: WinMLBuildConfig) -> None: preloaded_hf_config=preloaded_hf_config, output_dir=output_dir, use_cache=use_cache, - device=device, - ep=ep_value, + device=runtime_device, + ep=runtime_ep_value, precision=precision, rebuild=rebuild, submodel=submodel, @@ -1198,8 +1202,8 @@ def _patch_device(cfg: WinMLBuildConfig) -> None: configs=configs, output_dir=resolved_dir, rebuild=rebuild, - ep=ep_value, - device=device, + ep=runtime_ep_value, + device=runtime_device, allow_unsupported_nodes=allow_unsupported_nodes, ) @@ -1344,9 +1348,10 @@ def _patch_device(cfg: WinMLBuildConfig) -> None: model, task=component_task, trust_remote_code=trust_remote_code, - device=device, + device=runtime_device, precision=precision, - ep=ep_value, + ep=runtime_ep_value, + export_policy_target=(request_device, request_ep_value), shape_config=shape_overrides, override={"export": export_overrides} if export_overrides else None, ) @@ -1407,8 +1412,8 @@ def _patch_device(cfg: WinMLBuildConfig) -> None: resolved_dir=resolved_dir, rebuild=rebuild, cache_key=name, - ep=ep_value, - device=device, + ep=runtime_ep_value, + device=runtime_device, extra_kwargs=dict(extra_kwargs), preloaded_hf_config=preloaded_hf_config, ) @@ -1425,8 +1430,8 @@ def _patch_device(cfg: WinMLBuildConfig) -> None: resolved_dir=resolved_dir, rebuild=rebuild, cache_key=cache_key, - ep=ep_value, - device=device, + ep=runtime_ep_value, + device=runtime_device, extra_kwargs=extra_kwargs, preloaded_hf_config=preloaded_hf_config, ) diff --git a/src/winml/modelkit/commands/export.py b/src/winml/modelkit/commands/export.py index 2a2f5a8b0..12a9cc3ff 100644 --- a/src/winml/modelkit/commands/export.py +++ b/src/winml/modelkit/commands/export.py @@ -270,7 +270,12 @@ def export( if not cli_utils.is_cli_provided(ctx, "dynamo") and "dynamo" in ec: dynamo = ec["dynamo"] - from ..export import InputTensorSpec, OutputTensorSpec, WinMLExportConfig + from ..export import ( + InputTensorSpec, + OutputTensorSpec, + WinMLExportConfig, + resolve_export_compatibility, + ) from ..export import export_pytorch as export_onnx from ..loader import load_hf_model @@ -430,6 +435,8 @@ def _run_component_export(component_task: str | None, out_path: Path) -> None: try: cfg = WinMLExportConfig.from_dict(config_kwargs) + if not cfg.compatibility: + cfg.compatibility = resolve_export_compatibility() except Exception as e: console.print(f"[bold red]Configuration error:[/bold red] {e}") logger.exception("Failed to create export config") diff --git a/src/winml/modelkit/commands/perf.py b/src/winml/modelkit/commands/perf.py index 6166b7634..13f84d9ac 100644 --- a/src/winml/modelkit/commands/perf.py +++ b/src/winml/modelkit/commands/perf.py @@ -987,6 +987,8 @@ def _load_model(self) -> None: "task": resolved_task, "config": override, "ep_device": self._ep_device, + "device": self.config.device, + "ep": self.config.ep, "precision": self.config.precision, "provider_options": self.config.ep_options, "use_cache": use_cache, @@ -1346,8 +1348,10 @@ def _perf_modules( from ..session import EPDeviceTarget, WinMLEPRegistry, resolve_device from .build import _instantiate_parent_model + request_device = (device or "auto").lower() + request_ep = ep resolved_target = resolve_device( - EPDeviceTarget(ep=ep or "auto", device=device or "auto", source=ep_source) + EPDeviceTarget(ep=request_ep or "auto", device=request_device, source=ep_source) ) resolved_ep_device = WinMLEPRegistry.instance().auto_device(resolved_target) resolved_device = resolved_target.device @@ -1363,6 +1367,7 @@ def _perf_modules( device=resolved_device, precision=precision, ep=ep, + export_policy_target=(request_device, request_ep), ) except SubmoduleClassNotFoundError as e: # User-error: --module pattern didn't match. List what's available so @@ -2669,7 +2674,6 @@ def perf( except Exception as e: raise click.ClickException(f"Failed to resolve Hub-hosted ONNX path {model!r}: {e}") from e model = hf_model - # AC 11 (mockup spec): --top-k requires --op-tracing. Outside the # op-tracing section the flag is meaningless, so reject it explicitly # rather than silently ignoring a user's intent. diff --git a/src/winml/modelkit/config/build.py b/src/winml/modelkit/config/build.py index c622ab78a..e663e367f 100644 --- a/src/winml/modelkit/config/build.py +++ b/src/winml/modelkit/config/build.py @@ -58,6 +58,7 @@ WinMLExportConfig, _resolve_export_config_from_specs, ) +from ..export.policy import export_policy_targets_for_request, resolve_export_compatibility from ..loader.config import WinMLLoaderConfig, resolve_loader_config from ..optim.config import WinMLOptimizationConfig from ..quant.config import WinMLQuantizationConfig @@ -71,7 +72,7 @@ if TYPE_CHECKING: - from collections.abc import Mapping + from collections.abc import Mapping, Sequence import torch from torch import nn @@ -79,6 +80,8 @@ from ..eval.config import WinMLEvaluationConfig # noqa: TC004 from ..utils.constants import EPNameOrAlias +ExportPolicyTargetRequest = tuple[str | None, str | None] + __all__ = [ "WinMLBuildConfig", "generate_build_config", @@ -451,6 +454,34 @@ def _apply_target_policy( ) +def _is_explicit_export_policy_target(*, device: str | None, ep: str | None) -> bool: + """Return whether the request named a specific EP/device export target.""" + return (ep is not None and ep.lower() != "auto") or ( + device is not None and device.lower() != "auto" + ) + + +def apply_export_compatibility_policy( + config: WinMLBuildConfig | Sequence[WinMLBuildConfig], + *, + device: str | None = "auto", + ep: str | None = None, +) -> None: + """Populate export compatibility when the config has an export stage.""" + export_policy_targets = export_policy_targets_for_request( + ep=ep, + device=device, + target_was_explicit=_is_explicit_export_policy_target(device=device, ep=ep), + ) + configs = (config,) if isinstance(config, WinMLBuildConfig) else config + for cfg in configs: + if cfg.export is None: + continue + if cfg.export.compatibility: + continue + cfg.export.compatibility = resolve_export_compatibility(export_policy_targets) + + def resolve_quant_compile_config( *, device: str = "auto", @@ -794,6 +825,7 @@ def generate_hf_build_config( precision: str = "auto", trust_remote_code: bool = False, ep: str | None = None, + export_policy_target: ExportPolicyTargetRequest | None = None, policy_overrides_config: bool = False, no_compile: bool = False, ) -> WinMLBuildConfig: ... @@ -814,6 +846,7 @@ def generate_hf_build_config( precision: str = "auto", trust_remote_code: bool = False, ep: str | None = None, + export_policy_target: ExportPolicyTargetRequest | None = None, policy_overrides_config: bool = False, no_compile: bool = False, ) -> list[WinMLBuildConfig]: ... @@ -838,6 +871,7 @@ def generate_hf_build_config( precision: str = "auto", trust_remote_code: bool = False, ep: str | None = None, + export_policy_target: ExportPolicyTargetRequest | None = None, policy_overrides_config: bool = False, no_compile: bool = False, ) -> WinMLBuildConfig | list[WinMLBuildConfig]: ... @@ -857,6 +891,7 @@ def generate_hf_build_config( precision: str = "auto", trust_remote_code: bool = False, ep: str | None = None, + export_policy_target: ExportPolicyTargetRequest | None = None, policy_overrides_config: bool = False, no_compile: bool = False, ) -> WinMLBuildConfig | list[WinMLBuildConfig]: @@ -897,6 +932,9 @@ class name. Uses torchinfo to discover submodules and infer "int16", or "w{x}a{y}" e.g. "w8a16"). trust_remote_code: Allow running custom code from model repository. ep: Explicit execution provider override. + export_policy_target: Optional ``(device, ep)`` request used only for + export compatibility resolution. When omitted, the build target is + used for both quant/compile policy and export compatibility. policy_overrides_config: Apply device/precision/EP policy after ``override``. CLI callers set this only when a target option was explicitly supplied; otherwise sparse config values remain higher @@ -1070,6 +1108,11 @@ class name. Uses torchinfo to discover submodules and infer if no_compile: parent_config.compile = None + # Apply export compatibility policy so parent_config.export.compatibility is populated + # (used for serialization/cache-key participation and inheritance by submodules). + policy_device, policy_ep = export_policy_target or (device, ep) + apply_export_compatibility_policy(parent_config, device=policy_device, ep=policy_ep) + # ========================================================================= # STEP 5: Specialize for submodules if requested # ========================================================================= @@ -1130,6 +1173,7 @@ def generate_build_config( precision: str = "auto", trust_remote_code: bool = False, ep: str | None = None, + export_policy_target: ExportPolicyTargetRequest | None = None, onnx_path: str | Path | None = None, ) -> WinMLBuildConfig: ... @@ -1149,6 +1193,7 @@ def generate_build_config( precision: str = "auto", trust_remote_code: bool = False, ep: str | None = None, + export_policy_target: ExportPolicyTargetRequest | None = None, onnx_path: str | Path | None = None, ) -> list[WinMLBuildConfig]: ... @@ -1167,6 +1212,7 @@ def generate_build_config( precision: str = "auto", trust_remote_code: bool = False, ep: str | None = None, + export_policy_target: ExportPolicyTargetRequest | None = None, onnx_path: str | Path | None = None, ) -> WinMLBuildConfig | list[WinMLBuildConfig]: """Generate WinMLBuildConfig by orchestrating existing modules. @@ -1191,6 +1237,8 @@ class name (HF path only). "int16", or "w{x}a{y}" e.g. "w8a16"). trust_remote_code: Allow running custom code from model repository. ep: Explicit execution provider override. + export_policy_target: Optional ``(device, ep)`` request used only for + export compatibility resolution on HuggingFace exports. onnx_path: Path to a pre-exported ONNX file (Scenario D). Returns: @@ -1223,6 +1271,7 @@ class name (HF path only). precision=precision, trust_remote_code=trust_remote_code, ep=ep, + export_policy_target=export_policy_target, policy_overrides_config=True, ) @@ -1297,6 +1346,11 @@ def _input_name(i: int) -> str: if parent_config.export is not None else WinMLExportConfig().dynamo ), + compatibility=( + copy.deepcopy(parent_config.export.compatibility) + if parent_config.export is not None + else WinMLExportConfig().compatibility + ), # opset_version and batch_size use dataclass defaults from WinMLExportConfig ), optim=copy.deepcopy(parent_config.optim), @@ -1382,6 +1436,11 @@ def _merge_export_config( if override.hierarchy_tag_format != defaults.hierarchy_tag_format else base.hierarchy_tag_format ), + compatibility=( + copy.deepcopy(override.compatibility) + if override.compatibility + else copy.deepcopy(base.compatibility) + ), ) diff --git a/src/winml/modelkit/eval/evaluate.py b/src/winml/modelkit/eval/evaluate.py index bc2d0a861..635d034b8 100644 --- a/src/winml/modelkit/eval/evaluate.py +++ b/src/winml/modelkit/eval/evaluate.py @@ -335,6 +335,8 @@ def _load_model( config.model_id, ep_device, task=config.task, + device=config.device, + ep=config.ep, precision=config.precision, allow_unsupported_nodes=config.allow_unsupported_nodes, config=build_override, diff --git a/src/winml/modelkit/export/__init__.py b/src/winml/modelkit/export/__init__.py index f935f24d8..c0f41131e 100644 --- a/src/winml/modelkit/export/__init__.py +++ b/src/winml/modelkit/export/__init__.py @@ -19,6 +19,13 @@ WinMLExportConfig, resolve_export_config, ) +from .policy import ( + ExportCompatibilityConfig, + ExportCompatibilityRule, + ExportPolicyTarget, + export_policy_targets_for_request, + resolve_export_compatibility, +) # Static type re-exports for the names exposed by ``__getattr__`` below. @@ -40,15 +47,20 @@ __version__ = "2.1.0" __all__ = [ + "ExportCompatibilityConfig", + "ExportCompatibilityRule", + "ExportPolicyTarget", "InputTensorSpec", "MaxLengthTextInputGenerator", "ONNXConfigNotFoundError", "OutputTensorSpec", "WinMLExportConfig", "export_onnx", + "export_policy_targets_for_request", "export_pytorch", "generate_dummy_inputs", "register_onnx_overwrite", + "resolve_export_compatibility", "resolve_export_config", "resolve_io_specs", ] diff --git a/src/winml/modelkit/export/compatibility_rules.json b/src/winml/modelkit/export/compatibility_rules.json new file mode 100644 index 000000000..d8bdb575e --- /dev/null +++ b/src/winml/modelkit/export/compatibility_rules.json @@ -0,0 +1,15 @@ +{ + "schema_version": 1, + "rules": [ + { + "match": { + "ep": null, + "device": null + }, + "export": { + "transformers_attention": "eager" + }, + "reason": "Transformers SDPA-exported attention guard paths are not broadly portable." + } + ] +} diff --git a/src/winml/modelkit/export/config.py b/src/winml/modelkit/export/config.py index 8af7d308f..7649760ae 100644 --- a/src/winml/modelkit/export/config.py +++ b/src/winml/modelkit/export/config.py @@ -13,11 +13,12 @@ from __future__ import annotations import logging -from dataclasses import InitVar, dataclass +from dataclasses import InitVar, dataclass, field from typing import TYPE_CHECKING, Any, Literal # InputTensorSpec and OutputTensorSpec live in modelkit.onnx.io (canonical home). from ..onnx import InputTensorSpec, OutputTensorSpec +from .policy import ExportCompatibilityConfig if TYPE_CHECKING: @@ -179,6 +180,9 @@ class WinMLExportConfig: verbose: bool = False dynamo: bool = False # TorchScript exporter by default; True enables TorchDynamo + # Persisted export compatibility knobs (resolved from policy) + compatibility: ExportCompatibilityConfig = field(default_factory=ExportCompatibilityConfig) + # Phase 2: Hierarchy Preservation Options enable_hierarchy_tags: bool = True # Enable HTP hierarchy tagging by default clean_onnx: bool = False # Remove hierarchy tags for deployment @@ -214,6 +218,10 @@ def __post_init__( # Convert legacy output_names to output_tensors self.output_tensors = [OutputTensorSpec(name=name) for name in output_names_] + # Convert compatibility dicts (from deserialized data) to typed config + if isinstance(self.compatibility, dict): + self.compatibility = ExportCompatibilityConfig.from_dict(self.compatibility) + if self.batch_size <= 0: raise ValueError(f"batch_size must be positive, got {self.batch_size}") @@ -364,6 +372,9 @@ def to_dict(self) -> dict[str, Any]: if self.dynamic_axes: result["dynamic_axes"] = self.dynamic_axes + if self.compatibility: + result["compatibility"] = self.compatibility.to_dict() + return result @classmethod @@ -407,6 +418,11 @@ def from_dict(cls, data: dict[str, Any]) -> WinMLExportConfig: enable_hierarchy_tags=data.get("enable_hierarchy_tags", True), clean_onnx=data.get("clean_onnx", False), hierarchy_tag_format=data.get("hierarchy_tag_format", "full"), + compatibility=( + ExportCompatibilityConfig.from_dict(data["compatibility"]) + if "compatibility" in data + else ExportCompatibilityConfig() + ), ) diff --git a/src/winml/modelkit/export/htp/exporter.py b/src/winml/modelkit/export/htp/exporter.py index 90fe6e0eb..eb49a24c7 100644 --- a/src/winml/modelkit/export/htp/exporter.py +++ b/src/winml/modelkit/export/htp/exporter.py @@ -38,6 +38,7 @@ create_node_tagger_from_hierarchy, ) from ...core.onnx_utils import infer_output_names +from ...transformers_compat import use_eager_attention_for_export from .base_writer import ExportStep from .hierarchy import TracingHierarchyBuilder from .monitor import HTPExportMonitor @@ -595,7 +596,10 @@ def _convert_model_to_onnx( if export_config.dynamic_axes: onnx_kwargs["dynamic_axes"] = export_config.dynamic_axes - with self._get_optimum_patcher(model, task): + with ( + self._get_optimum_patcher(model, task), + self._export_compatibility_context(model, export_config), + ): # Models can override input binding by implementing # get_export_args(inputs) → tuple of positional args. # Default: pass inputs dict as kwargs. @@ -606,6 +610,16 @@ def _convert_model_to_onnx( else: torch.onnx.export(model, (), output_path, kwargs=inputs, **onnx_kwargs) + @staticmethod + def _export_compatibility_context( + model: nn.Module, + export_config: WinMLExportConfig, + ) -> contextlib.AbstractContextManager[None]: + """Return the export-time compatibility context requested by policy.""" + if export_config.compatibility.transformers_attention == "eager": + return use_eager_attention_for_export(model) + return contextlib.nullcontext() + @staticmethod def _resolve_keyword_input_names( model: nn.Module, diff --git a/src/winml/modelkit/export/policy.py b/src/winml/modelkit/export/policy.py new file mode 100644 index 000000000..2e54c4c84 --- /dev/null +++ b/src/winml/modelkit/export/policy.py @@ -0,0 +1,231 @@ +# ------------------------------------------------------------------------- +# Copyright (c) Microsoft Corporation. All rights reserved. +# Licensed under the MIT License. +# -------------------------------------------------------------------------- + +from __future__ import annotations + +import json +from dataclasses import dataclass +from functools import cache +from importlib import resources +from typing import TYPE_CHECKING, Any + +from ..utils.constants import EP_SUPPORTED_DEVICES, normalize_ep_name + + +if TYPE_CHECKING: + from collections.abc import Sequence + + +@dataclass(frozen=True) +class ExportCompatibilityConfig: + """Resolved export-time compatibility knobs.""" + + transformers_attention: str | None = None + + def __bool__(self) -> bool: # pragma: no cover - trivial + return self.transformers_attention is not None + + def to_dict(self) -> dict[str, str]: + """Serialize resolved compatibility knobs to a dict.""" + result: dict[str, str] = {} + if self.transformers_attention is not None: + result["transformers_attention"] = self.transformers_attention + return result + + @classmethod + def from_dict( + cls, + data: dict[str, Any] | None, + ) -> ExportCompatibilityConfig: + """Deserialize compatibility config from a dict (or None).""" + if data is None: + return cls() + if not isinstance(data, dict): + raise TypeError(f"export.compatibility must be an object, got {type(data).__name__}") + unknown = set(data) - {"transformers_attention"} + if unknown: + raise ValueError(f"Unknown export.compatibility field(s): {sorted(unknown)}") + attention = data.get("transformers_attention") + if attention is not None and attention != "eager": + raise ValueError( + "export.compatibility.transformers_attention must be 'eager' or null, " + f"got {attention!r}" + ) + return cls(transformers_attention=attention) + + +@dataclass(frozen=True) +class ExportPolicyTarget: + """One EP/device target used by export compatibility policy.""" + + ep: str + device: str + + def __post_init__(self) -> None: + normalized_ep = normalize_ep_name(self.ep) + object.__setattr__(self, "ep", normalized_ep if normalized_ep is not None else self.ep) + object.__setattr__(self, "device", self.device.lower()) + + +@dataclass(frozen=True) +class ExportCompatibilityRule: + """One EP/device compatibility rule.""" + + ep: str | None + device: str | None + compatibility: ExportCompatibilityConfig + reason: str + + def matches(self, target: ExportPolicyTarget) -> bool: + """Return True if this rule applies to the given target.""" + return (self.ep is None or target.ep == normalize_ep_name(self.ep)) and ( + self.device is None or target.device == self.device + ) + + +_RULES_RESOURCE = "compatibility_rules.json" +_SUPPORTED_POLICY_DEVICES_BY_EP: dict[str, frozenset[str]] = { + ep: frozenset(devices) for ep, devices in EP_SUPPORTED_DEVICES.items() +} +_SUPPORTED_POLICY_DEVICES = frozenset( + device for devices in _SUPPORTED_POLICY_DEVICES_BY_EP.values() for device in devices +) + + +def export_policy_targets_for_request( + *, + ep: str | None, + device: str | None, + target_was_explicit: bool, +) -> tuple[ExportPolicyTarget, ...] | None: + """Return explicit policy targets, or None for the portable catalog default.""" + if not target_was_explicit: + return None + + from ..session import EPDeviceTarget, resolve_device + + resolved = resolve_device(EPDeviceTarget(ep=ep or "auto", device=(device or "auto").lower())) + return (ExportPolicyTarget(ep=resolved.ep, device=resolved.device),) + + +def load_export_compatibility_rules() -> tuple[ExportCompatibilityRule, ...]: + """Load built-in export compatibility rules from package JSON.""" + return _load_export_compatibility_rules() + + +@cache +def _load_export_compatibility_rules() -> tuple[ExportCompatibilityRule, ...]: + data = json.loads(resources.files(__package__).joinpath(_RULES_RESOURCE).read_text()) + if data.get("schema_version") != 1: + raise ValueError( + f"{_RULES_RESOURCE} schema_version must be 1, got {data.get('schema_version')!r}" + ) + rules = data.get("rules") + if not isinstance(rules, list): + raise TypeError(f"{_RULES_RESOURCE} must contain a 'rules' array") + return tuple(_rule_from_dict(rule, index=index) for index, rule in enumerate(rules)) + + +def resolve_export_compatibility( + targets: Sequence[ExportPolicyTarget] | None = None, + *, + rules: Sequence[ExportCompatibilityRule] | None = None, +) -> ExportCompatibilityConfig: + """Resolve export compatibility for explicit targets or the portable catalog.""" + rules = load_export_compatibility_rules() if rules is None else rules + resolved_targets = _catalog_targets() if targets is None else tuple(targets) + + transformers_attention: str | None = None + transformers_attention_source: str | None = None + + for target in resolved_targets: + for rule in rules: + if not rule.matches(target): + continue + incoming = rule.compatibility.transformers_attention + if incoming is None: + continue + if transformers_attention is None: + transformers_attention = incoming + transformers_attention_source = f"{rule.ep or '*'}/{rule.device or '*'}" + elif transformers_attention != incoming: + raise ValueError( + "Conflicting export compatibility for transformers_attention: " + f"{transformers_attention!r} from {transformers_attention_source} vs " + f"{incoming!r} from {rule.ep or '*'}/{rule.device or '*'}" + ) + + return ExportCompatibilityConfig(transformers_attention=transformers_attention) + + +def _catalog_targets() -> tuple[ExportPolicyTarget, ...]: + from ..session.ep_device import EP_DEVICE_SPECS + + return tuple(ExportPolicyTarget(ep=spec.ep, device=spec.device) for spec in EP_DEVICE_SPECS) + + +def _rule_from_dict(data: object, *, index: int) -> ExportCompatibilityRule: + if not isinstance(data, dict): + raise TypeError(f"{_RULES_RESOURCE} rules[{index}] must be an object") + unknown = set(data) - {"match", "export", "reason"} + if unknown: + raise ValueError( + f"{_RULES_RESOURCE} rules[{index}] has unknown field(s): {sorted(unknown)}" + ) + + match = data.get("match") + if not isinstance(match, dict): + raise TypeError(f"{_RULES_RESOURCE} rules[{index}].match must be an object") + match_unknown = set(match) - {"ep", "device"} + if match_unknown: + raise ValueError( + f"{_RULES_RESOURCE} rules[{index}].match has unknown field(s): {sorted(match_unknown)}" + ) + ep = match.get("ep") + normalized_ep: str | None = None + if ep is not None and (not isinstance(ep, str) or not ep): + raise ValueError( + f"{_RULES_RESOURCE} rules[{index}].match.ep must be a non-empty string or null" + ) + if isinstance(ep, str): + normalized_ep = normalize_ep_name(ep) + if normalized_ep not in _SUPPORTED_POLICY_DEVICES_BY_EP: + raise ValueError( + f"{_RULES_RESOURCE} rules[{index}].match.ep must be a supported EP " + f"or alias, got {ep!r}" + ) + device = match.get("device") + if device is not None and (not isinstance(device, str) or not device): + raise ValueError( + f"{_RULES_RESOURCE} rules[{index}].match.device must be a non-empty string or null" + ) + normalized_device = device.lower() if isinstance(device, str) else None + if normalized_device is not None and normalized_device not in _SUPPORTED_POLICY_DEVICES: + raise ValueError( + f"{_RULES_RESOURCE} rules[{index}].match.device must be one of " + f"{sorted(_SUPPORTED_POLICY_DEVICES)}, got {device!r}" + ) + if ( + normalized_ep is not None + and normalized_device is not None + and normalized_device not in _SUPPORTED_POLICY_DEVICES_BY_EP[normalized_ep] + ): + raise ValueError( + f"{_RULES_RESOURCE} rules[{index}].match.ep {normalized_ep!r} " + f"does not support device {normalized_device!r}" + ) + + export = data.get("export") + compatibility = ExportCompatibilityConfig.from_dict(export) + reason = data.get("reason") + if not isinstance(reason, str) or not reason: + raise ValueError(f"{_RULES_RESOURCE} rules[{index}].reason must be a non-empty string") + + return ExportCompatibilityRule( + ep=normalized_ep, + device=normalized_device, + compatibility=compatibility, + reason=reason, + ) diff --git a/src/winml/modelkit/models/auto.py b/src/winml/modelkit/models/auto.py index 922aa1254..36131134d 100644 --- a/src/winml/modelkit/models/auto.py +++ b/src/winml/modelkit/models/auto.py @@ -370,6 +370,8 @@ def from_pretrained( model_input = resolve_model_input(str(model_id_or_path)) model_id = model_input.local_path or model_input.raw logger.info("Loading WinML model from: %s", model_id) + request_device = (device or "auto").lower() + request_ep = ep # Resolve a concrete target before every dispatch path, including # composites. Explicit incompatible requests intentionally propagate. @@ -448,7 +450,8 @@ def from_pretrained( return composite_cls.from_pretrained( model_id, task, - device=ep_device.device.device_type.lower(), + device=request_device, + ep=request_ep, ep_device=ep_device, use_cache=use_cache, force_rebuild=force_rebuild, @@ -474,6 +477,7 @@ def from_pretrained( task=task, config=config, ep_device=ep_device, + device=request_device, ep=ep, precision=precision, cache_dir=cache_dir, @@ -533,14 +537,16 @@ def _build_pretrained_artifact( model_input = resolve_model_input(str(model_id_or_path)) model_id = model_input.local_path or model_input.raw + request_device = (device or "auto").lower() + request_ep = ep if ep_device is None: from ..session import EPDeviceTarget, WinMLEPRegistry, resolve_device - target = resolve_device( - EPDeviceTarget(ep=ep or "auto", device=(device or "auto").lower()) - ) + target = resolve_device(EPDeviceTarget(ep=request_ep or "auto", device=request_device)) ep_device = WinMLEPRegistry.instance().auto_device(target) + runtime_device = ep_device.device.device_type.lower() + runtime_ep = _resolved_ep_short_name(ep_device) from ..config import generate_hf_build_config @@ -549,9 +555,10 @@ def _build_pretrained_artifact( task=task, override=config, shape_config=shape_config, - device=ep_device.device.device_type.lower(), + device=runtime_device, precision=precision, - ep=_resolved_ep_short_name(ep_device), + ep=runtime_ep, + export_policy_target=(request_device, request_ep), model_type=model_type, trust_remote_code=trust_remote_code, policy_overrides_config=True, diff --git a/src/winml/modelkit/models/winml/composite_model.py b/src/winml/modelkit/models/winml/composite_model.py index e1f6250b2..1455ba737 100644 --- a/src/winml/modelkit/models/winml/composite_model.py +++ b/src/winml/modelkit/models/winml/composite_model.py @@ -223,6 +223,8 @@ def from_pretrained( model_id, ep_device=ep_device, task=component_task, + device=device, + ep=ep, use_cache=use_cache, force_rebuild=force_rebuild, trust_remote_code=trust_remote_code, diff --git a/src/winml/modelkit/transformers_compat.py b/src/winml/modelkit/transformers_compat.py index 8dbb989e6..6f8e4f1d1 100644 --- a/src/winml/modelkit/transformers_compat.py +++ b/src/winml/modelkit/transformers_compat.py @@ -2,17 +2,20 @@ # Copyright (c) Microsoft Corporation. All rights reserved. # Licensed under the MIT License. # -------------------------------------------------------------------------- -"""Compat shim for optimum-onnx 0.1.0 against transformers 5.x, armed lazily. +"""Compatibility helpers for transformers-dependent export paths. + +The import hook below provides an optimum-onnx 0.1.0 shim against transformers +5.x, armed lazily. optimum-onnx 0.1.0 (last PyPI release as of 2026-04-30) hardcodes imports against transformers 4.x internals. This module re-injects those symbols so optimum-onnx's imports succeed on transformers 5.x. -Module load only inserts a ``sys.meta_path`` finder — it does NOT load -transformers. The finder calls :func:`install` the first time anything -imports ``optimum.*``; :func:`install` is idempotent. Lightweight commands -that never touch optimum (``winml sys``, ``winml --help``) pay zero -transformers cost. +Module load only defines lightweight helpers and inserts a ``sys.meta_path`` +finder — it does NOT load transformers. The finder calls :func:`install` the +first time anything imports ``optimum.*``; :func:`install` is idempotent. +Lightweight commands that never touch optimum (``winml sys``, ``winml --help``) +pay zero transformers cost. Drop this file (and the corresponding override in pyproject.toml) once optimum-onnx 0.2+ ships with transformers 5.x compatibility. @@ -20,16 +23,19 @@ from __future__ import annotations +import contextlib import sys from importlib.abc import Loader, MetaPathFinder from typing import TYPE_CHECKING, Any, cast if TYPE_CHECKING: - from collections.abc import Sequence + from collections.abc import Iterator, Sequence from importlib.machinery import ModuleSpec from types import ModuleType + import torch.nn as nn + _installed = False @@ -40,6 +46,100 @@ _MODEL_PATCHER_MODULE = "optimum.exporters.onnx.model_patcher" +@contextlib.contextmanager +def use_eager_attention_for_export(model: nn.Module) -> Iterator[None]: + """Temporarily prefer eager attention on HF-style module configs.""" + restored: list[tuple[int, Any, Any]] = [] + configs: dict[int, Any] = {} + children: dict[int, set[int]] = {} + seen_configs: set[int] = set() + + for module in model.modules(): + _collect_attention_configs(getattr(module, "config", None), configs, children, seen_configs) + + for config_id, config in configs.items(): + current = config._attn_implementation + restored.append((config_id, config, current)) + + for config in configs.values(): + if config._attn_implementation != "eager": + config._attn_implementation = "eager" + + try: + yield + finally: + for _config_id, config, previous in _parent_before_child(restored, children): + config._attn_implementation = previous + + +def _collect_attention_configs( + config: Any, + configs: dict[int, Any], + children: dict[int, set[int]], + seen_configs: set[int], +) -> None: + if config is None or id(config) in seen_configs: + return + seen_configs.add(id(config)) + + config = cast("Any", config) + config_id = id(config) + if hasattr(config, "_attn_implementation"): + configs[config_id] = config + + for child_config in _iter_sub_configs(config): + if hasattr(config, "_attn_implementation") and hasattr( + child_config, "_attn_implementation" + ): + children.setdefault(config_id, set()).add(id(child_config)) + _collect_attention_configs(child_config, configs, children, seen_configs) + + +def _iter_sub_configs(config: Any) -> Iterator[Any]: + sub_configs = getattr(config, "sub_configs", None) + if isinstance(sub_configs, dict): + for key, value in sub_configs.items(): + if isinstance(key, str): + child = getattr(config, key, None) + if child is not None: + yield child + if not isinstance(value, type): + yield value + elif isinstance(sub_configs, (list, tuple, set)): + yield from sub_configs + + +def _parent_before_child( + restored: list[tuple[int, Any, Any]], + children: dict[int, set[int]], +) -> list[tuple[int, Any, Any]]: + parents: dict[int, set[int]] = {} + restored_ids = {config_id for config_id, _config, _previous in restored} + for parent_id, child_ids in children.items(): + if parent_id not in restored_ids: + continue + for child_id in child_ids: + if child_id in restored_ids: + parents.setdefault(child_id, set()).add(parent_id) + + depths: dict[int, int] = {} + + def depth(config_id: int, visiting: set[int]) -> int: + if config_id in depths: + return depths[config_id] + if config_id in visiting: + return 0 + visiting.add(config_id) + config_depth = 0 + if config_id in parents: + config_depth = 1 + max(depth(parent_id, visiting) for parent_id in parents[config_id]) + visiting.remove(config_id) + depths[config_id] = config_depth + return config_depth + + return sorted(restored, key=lambda item: depth(item[0], set())) + + def install() -> None: """Apply the transformers 5.x ↔ optimum-onnx 0.1.0 shim. Idempotent.""" global _installed diff --git a/tests/unit/commands/test_build.py b/tests/unit/commands/test_build.py index af27b01d6..37f6bde79 100644 --- a/tests/unit/commands/test_build.py +++ b/tests/unit/commands/test_build.py @@ -765,6 +765,26 @@ def test_basic_build( assert result.exit_code == 0, f"Build failed: {result.output}" assert mock_build_api.called + def test_loaded_config_gets_default_export_compatibility( + self, + runner: CliRunner, + sample_config_file: Path, + mock_build_api: MagicMock, + tmp_path: Path, + ) -> None: + from winml.modelkit.commands.build import build + + result = runner.invoke( + build, + ["-c", str(sample_config_file), "-m", "test-model", "-o", str(tmp_path / "out")], + obj={"debug": False}, + ) + assert result.exit_code == 0, result.output + + config = mock_build_api.call_args.kwargs["config"] + assert config.export is not None + assert config.export.compatibility.transformers_attention == "eager" + def test_model_id_passed( self, runner: CliRunner, @@ -1107,6 +1127,70 @@ def test_device_flag_passed( call_kwargs = mock_build_api.call_args.kwargs assert call_kwargs["device"] == "npu" + def test_auto_generated_config_leaves_export_policy_to_config_generation( + self, + runner: CliRunner, + mock_build_api: MagicMock, + tmp_path: Path, + ) -> None: + from winml.modelkit.commands.build import build + from winml.modelkit.config import WinMLBuildConfig + + fake_cfg = WinMLBuildConfig.from_dict( + { + "loader": {"task": "image-classification"}, + "export": {"opset_version": 17, "batch_size": 1}, + "optim": {}, + "quant": None, + "compile": None, + } + ) + with patch( + "winml.modelkit.config.generate_build_config", return_value=fake_cfg + ) as mock_gen: + result = runner.invoke( + build, + ["-m", "microsoft/resnet-50", "-o", str(tmp_path), "--ep", "qnn"], + obj={"debug": False}, + ) + + assert result.exit_code == 0, result.output + assert mock_gen.call_args.kwargs["device"] == "auto" + assert mock_gen.call_args.kwargs["ep"] == "qnn" + assert mock_gen.call_args.kwargs["export_policy_target"] == ("auto", "qnn") + + def test_auto_generated_config_uses_portable_export_policy_when_no_target_supplied( + self, + runner: CliRunner, + mock_build_api: MagicMock, + tmp_path: Path, + ) -> None: + from winml.modelkit.commands.build import build + from winml.modelkit.config import WinMLBuildConfig + + fake_cfg = WinMLBuildConfig.from_dict( + { + "loader": {"task": "image-classification"}, + "export": {"opset_version": 17, "batch_size": 1}, + "optim": {}, + "quant": None, + "compile": None, + } + ) + with patch( + "winml.modelkit.config.generate_build_config", return_value=fake_cfg + ) as mock_gen: + result = runner.invoke( + build, + ["-m", "microsoft/resnet-50", "-o", str(tmp_path)], + obj={"debug": False}, + ) + + assert result.exit_code == 0, result.output + assert mock_gen.call_args.kwargs["device"] == "npu" + assert mock_gen.call_args.kwargs["ep"] == "QNNExecutionProvider" + assert mock_gen.call_args.kwargs["export_policy_target"] == ("auto", None) + def test_input_specs_patches_config_file_inputs( self, runner: CliRunner, @@ -1967,6 +2051,33 @@ def test_ep_forwarded_to_generate_build_config( assert result.exit_code == 0, result.output assert mock_gen.call_args.kwargs["ep"] == "openvino" + def test_auto_config_uses_resolved_runtime_target_for_build_policy( + self, tmp_path: Path, mock_run_single_build: MagicMock + ) -> None: + """Auto-generated configs use runtime target while export policy keeps request.""" + fake_cfg = MagicMock() + fake_cfg.compile = None + fake_cfg.validate.return_value = None + fake_cfg.loader = MagicMock() + fake_cfg.loader.task = "image-classification" + resolved_target = EPDeviceTarget(ep="DmlExecutionProvider", device="gpu") + + with ( + patch("winml.modelkit.session.resolve_device", return_value=resolved_target), + patch("winml.modelkit.config.generate_build_config", return_value=fake_cfg) as mock_gen, + patch( + "winml.modelkit.commands.build._validate_loader_tasks_for_model", + return_value=None, + ), + ): + result = _invoke(["-m", "microsoft/resnet-50", "-o", str(tmp_path / "out")]) + + assert result.exit_code == 0, result.output + kwargs = mock_gen.call_args.kwargs + assert kwargs["device"] == "gpu" + assert kwargs["ep"] == "DmlExecutionProvider" + assert kwargs["export_policy_target"] == ("auto", None) + def test_export_overrides_forwarded_to_generate_build_config( self, tmp_path: Path, mock_run_single_build: MagicMock ): @@ -2344,6 +2455,17 @@ def test_composite_builds_flat_with_cache_key( if "task" in call.kwargs and call.kwargs["task"] is not None } assert component_tasks == set(components.values()) + component_config_targets = { + call.kwargs["task"]: ( + call.kwargs["device"], + call.kwargs["ep"], + call.kwargs["export_policy_target"], + ) + for call in mock_gen_cfg.call_args_list + if "task" in call.kwargs and call.kwargs["task"] is not None + } + expected_target = ("npu", "QNNExecutionProvider", ("auto", None)) + assert component_config_targets == dict.fromkeys(components.values(), expected_target) def test_composite_autogen_passes_task_none_to_resolver( self, diff --git a/tests/unit/commands/test_export.py b/tests/unit/commands/test_export.py index 9218c084c..97a40868a 100644 --- a/tests/unit/commands/test_export.py +++ b/tests/unit/commands/test_export.py @@ -203,6 +203,36 @@ def test_export_calls_export_onnx( assert "model_id" in call_kwargs assert call_kwargs["model_id"] == "test-model" + def test_export_resolves_portable_compatibility_by_default( + self, + runner: CliRunner, + mock_export_onnx: MagicMock, + mock_load_hf_model: MagicMock, + tmp_path: Path, + ) -> None: + """Standalone exports should get the same portable policy as build configs.""" + from winml.modelkit.commands.export import export + from winml.modelkit.export import WinMLExportConfig + + with patch( + "winml.modelkit.export.resolve_export_config", + return_value=(WinMLExportConfig(), MagicMock()), + ): + result = runner.invoke( + export, + [ + "--model", + "test-model", + "--output", + str(tmp_path / "model.onnx"), + ], + obj={"debug": False}, + ) + + assert result.exit_code == 0, result.output + cfg = mock_export_onnx.call_args.kwargs["export_config"] + assert cfg.compatibility.transformers_attention == "eager" + def test_export_passes_verbose_flag( self, runner: CliRunner, diff --git a/tests/unit/commands/test_perf_cli.py b/tests/unit/commands/test_perf_cli.py index bc4d38c0d..cc2c07c1c 100644 --- a/tests/unit/commands/test_perf_cli.py +++ b/tests/unit/commands/test_perf_cli.py @@ -186,6 +186,37 @@ def test_path_is_under_user_home(self) -> None: class TestPerfUnifiedPipeline: """Test that both ONNX and HF models go through PerfBenchmark._load_model.""" + def test_load_model_does_not_forward_export_policy_details( + self, monkeypatch: pytest.MonkeyPatch + ) -> None: + """Pre-resolved runtime EPs should not force target-specific export policy.""" + from winml.modelkit.models import WinMLAutoModel + + benchmark = PerfBenchmark( + BenchmarkConfig( + model_id="microsoft/resnet-50", + task="image-classification", + ) + ) + fake_ep_device = MagicMock() + benchmark._ep_device = fake_ep_device + benchmark._resolved_device = "gpu" + benchmark._resolved_ep = "DmlExecutionProvider" + monkeypatch.setattr(benchmark, "_resolve_device_ep", lambda: None) + + received: dict[str, object] = {} + + def _from_pretrained(*args: object, **kwargs: object) -> MagicMock: + received["args"] = args + received.update(kwargs) + return MagicMock() + + monkeypatch.setattr(WinMLAutoModel, "from_pretrained", _from_pretrained) + + benchmark._load_model() + + assert received["ep_device"] is fake_ep_device + def test_close_releases_single_model_session(self) -> None: """Closing a benchmark resets the loaded model's native session.""" benchmark = PerfBenchmark(BenchmarkConfig(model_id="m")) diff --git a/tests/unit/commands/test_perf_module.py b/tests/unit/commands/test_perf_module.py index 2931c4f12..86adb3c9f 100644 --- a/tests/unit/commands/test_perf_module.py +++ b/tests/unit/commands/test_perf_module.py @@ -182,7 +182,7 @@ def test_device_and_ep_forwarded_through_module_path(self, tmp_path: Path) -> No fake_loader_cfg.task = "fill-mask" resolved_target = EPDeviceTarget( ep="QNNExecutionProvider", - device="npu", + device="gpu", source="pypi", ) resolved_ep_device = MagicMock(name="resolved_ep_device") @@ -244,13 +244,14 @@ def test_device_and_ep_forwarded_through_module_path(self, tmp_path: Path) -> No assert result.exit_code == 0, result.output gen_kwargs = mock_gen.call_args.kwargs - assert gen_kwargs["device"] == "npu" + assert gen_kwargs["device"] == "gpu" assert gen_kwargs["ep"] == "QNNExecutionProvider" + assert gen_kwargs["export_policy_target"] == ("npu", "qnn") assert gen_kwargs["precision"] == "auto" build_kwargs = mock_build.call_args.kwargs assert build_kwargs["ep"] == "QNNExecutionProvider" - assert build_kwargs["device"] == "npu" + assert build_kwargs["device"] == "gpu" fake_registry.auto_device.assert_called_once_with(resolved_target) session_kwargs = mock_session_cls.call_args.kwargs @@ -358,6 +359,70 @@ def test_running_model_path_in_module_result(self, tmp_path: Path) -> None: instance = report["instances"][0] assert instance["running_model_path"] == str(running_model_path) + def test_module_path_defaults_to_portable_policy_when_no_target_supplied( + self, tmp_path: Path + ) -> None: + fake_cfg = MagicMock() + fake_cfg.loader.model_type = "bert" + fake_cfg.loader.module_path = "encoder.layer.0" + + fake_build_result = MagicMock() + fake_build_result.final_onnx_path = tmp_path / "model.onnx" + + fake_session = MagicMock() + fake_session.perf.side_effect = RuntimeError("test-skip-benchmark") + fake_loader_cfg = MagicMock() + fake_loader_cfg.task = "fill-mask" + + with ( + patch( + "winml.modelkit.config.generate_hf_build_config", + return_value=[fake_cfg], + ) as mock_gen, + patch( + "winml.modelkit.loader.resolve_loader_config", + return_value=(fake_loader_cfg, MagicMock(), MagicMock(), MagicMock()), + ), + patch( + "winml.modelkit.commands.build._instantiate_parent_model", + return_value=MagicMock(), + ), + patch( + "winml.modelkit.build.build_hf_model", + return_value=fake_build_result, + ), + patch( + "winml.modelkit.session.WinMLSession", + return_value=fake_session, + ), + patch( + "winml.modelkit.commands.perf.generate_random_inputs", + return_value={}, + ), + ): + runner = CliRunner() + result = runner.invoke( + main, + [ + "perf", + "-m", + "fake/model", + "--module", + "BertLayer", + "--iterations", + "1", + "--warmup", + "0", + "-o", + str(tmp_path / "out.json"), + ], + ) + assert result.exit_code == 0, result.output + assert result.exit_code == 0, result.output + assert mock_gen.call_args.kwargs["device"] == "cpu" + assert mock_gen.call_args.kwargs["ep"] == "auto" + assert mock_gen.call_args.kwargs["export_policy_target"] == ("auto", None) + class TestPerfModuleMonitor: """--monitor must drive the live HW utilization chart in --module mode. diff --git a/tests/unit/config/test_build.py b/tests/unit/config/test_build.py index ac1345878..81bb3f653 100644 --- a/tests/unit/config/test_build.py +++ b/tests/unit/config/test_build.py @@ -43,6 +43,7 @@ WinMLExportConfig, resolve_io_specs, ) +from winml.modelkit.export.policy import ExportCompatibilityConfig from winml.modelkit.loader import WinMLLoaderConfig from winml.modelkit.optim import WinMLOptimizationConfig from winml.modelkit.quant import WinMLQuantizationConfig @@ -128,6 +129,187 @@ def mock_io_specs() -> dict: } +class TestGeneratedExportCompatibilityPolicy: + def test_generated_hf_config_uses_portable_policy_by_default( + self, + monkeypatch: pytest.MonkeyPatch, + mock_loader_config: WinMLLoaderConfig, + mock_hf_config: MagicMock, + mock_model_class: MagicMock, + mock_export_config: WinMLExportConfig, + ) -> None: + monkeypatch.setattr( + "winml.modelkit.config.build.resolve_loader_config", + lambda *args, **kwargs: ( + mock_loader_config, + mock_hf_config, + mock_model_class, + MagicMock(), + ), + ) + monkeypatch.setattr( + "winml.modelkit.config.build._resolve_export_config_from_specs", + lambda *args, **kwargs: mock_export_config, + ) + + cfg = generate_hf_build_config( + "local-model", + device="auto", + ep=None, + policy_overrides_config=True, + ) + + assert cfg.export is not None + assert cfg.export.compatibility.transformers_attention == "eager" + + def test_generated_hf_config_applies_global_export_policy_to_explicit_target( + self, + monkeypatch: pytest.MonkeyPatch, + mock_loader_config: WinMLLoaderConfig, + mock_hf_config: MagicMock, + mock_model_class: MagicMock, + mock_export_config: WinMLExportConfig, + ) -> None: + # MonkeyPatch fixture typing note. + monkeypatch.setattr( + "winml.modelkit.config.build.resolve_loader_config", + lambda *args, **kwargs: ( + mock_loader_config, + mock_hf_config, + mock_model_class, + MagicMock(), + ), + ) + monkeypatch.setattr( + "winml.modelkit.config.build._resolve_export_config_from_specs", + lambda *args, **kwargs: mock_export_config, + ) + + cfg = generate_hf_build_config( + "local-model", + device="gpu", + ep="DmlExecutionProvider", + policy_overrides_config=True, + ) + + assert cfg.export is not None + assert cfg.export.compatibility.transformers_attention == "eager" + + def test_generated_hf_config_can_split_build_and_export_policy_targets( + self, + monkeypatch: pytest.MonkeyPatch, + mock_loader_config: WinMLLoaderConfig, + mock_hf_config: MagicMock, + mock_model_class: MagicMock, + mock_export_config: WinMLExportConfig, + ) -> None: + target_policy_calls: list[tuple[str, str, str | None]] = [] + export_policy_calls: list[tuple[str | None, str | None]] = [] + + monkeypatch.setattr( + "winml.modelkit.config.build.resolve_loader_config", + lambda *args, **kwargs: ( + mock_loader_config, + mock_hf_config, + mock_model_class, + MagicMock(), + ), + ) + monkeypatch.setattr( + "winml.modelkit.config.build._resolve_export_config_from_specs", + lambda *args, **kwargs: mock_export_config, + ) + monkeypatch.setattr( + "winml.modelkit.config.build._apply_target_policy", + lambda config, *, device, precision, ep: target_policy_calls.append( + (device, precision, ep) + ), + ) + monkeypatch.setattr( + "winml.modelkit.config.build.apply_export_compatibility_policy", + lambda config, *, device, ep: export_policy_calls.append((device, ep)), + ) + + generate_hf_build_config( + "local-model", + device="gpu", + ep="DmlExecutionProvider", + export_policy_target=("auto", None), + policy_overrides_config=True, + ) + + assert target_policy_calls == [("gpu", "auto", "DmlExecutionProvider")] + assert export_policy_calls == [("auto", None)] + + def test_submodule_config_inherits_export_compatibility(self) -> None: + parent = WinMLBuildConfig( + loader=WinMLLoaderConfig(model_type="bert", task="fill-mask"), + export=WinMLExportConfig( + compatibility=ExportCompatibilityConfig(transformers_attention="eager") + ), + ) + sub_info = SubmoduleInfo( + class_name="Linear", + module_path="encoder.layer.0.output.dense", + input_shapes=[[1, 4]], + output_shapes=[[1, 4]], + input_dtypes=["float32"], + output_dtypes=["float32"], + input_names=["hidden_states"], + ) + + sub_cfg = _build_submodule_config(sub_info, parent) + + assert sub_cfg.export is not None + assert sub_cfg.export.compatibility.transformers_attention == "eager" + + +class TestLoadedConfigExportCompatibilityPolicy: + def test_apply_export_policy_populates_loaded_config_without_compatibility(self) -> None: + from winml.modelkit.config.build import apply_export_compatibility_policy + + cfg = WinMLBuildConfig(export=WinMLExportConfig()) + + apply_export_compatibility_policy(cfg) + + assert cfg.export is not None + assert cfg.export.compatibility.transformers_attention == "eager" + + def test_apply_export_policy_preserves_serialized_compatibility(self) -> None: + from winml.modelkit.config.build import apply_export_compatibility_policy + + cfg = WinMLBuildConfig( + export=WinMLExportConfig( + compatibility=ExportCompatibilityConfig(transformers_attention="eager") + ) + ) + + apply_export_compatibility_policy(cfg, device="gpu", ep="DmlExecutionProvider") + + assert cfg.export is not None + assert cfg.export.compatibility.transformers_attention == "eager" + + def test_serialized_empty_compatibility_receives_policy(self) -> None: + from winml.modelkit.config.build import apply_export_compatibility_policy + + cfg = WinMLBuildConfig.from_dict({"export": {"compatibility": {}}}) + apply_export_compatibility_policy(cfg) + + assert cfg.export is not None + assert cfg.export.compatibility.transformers_attention == "eager" + + def test_apply_export_policy_accepts_config_lists(self) -> None: + from winml.modelkit.config.build import apply_export_compatibility_policy + + cfgs = [WinMLBuildConfig(export=WinMLExportConfig()), WinMLBuildConfig(export=None)] + + apply_export_compatibility_policy(cfgs) + + assert cfgs[0].export is not None + assert cfgs[0].export.compatibility.transformers_attention == "eager" + assert cfgs[1].export is None + + # ============================================================================= # TestGetIoSpecsFromConfig - Unit tests for resolve_io_specs() # ============================================================================= @@ -270,6 +452,30 @@ def test_non_input_overrides_merged(self) -> None: assert [t.name for t in merged.export.input_tensors] == ["input_ids", "attention_mask"] +class TestExportCompatibilityBuildConfig: + def test_export_compatibility_changes_cache_key(self) -> None: + default_config = WinMLBuildConfig(export=WinMLExportConfig()) + eager_config = WinMLBuildConfig( + export=WinMLExportConfig( + compatibility=ExportCompatibilityConfig(transformers_attention="eager") + ) + ) + + assert default_config.generate_cache_key() != eager_config.generate_cache_key() + + def test_registered_export_merge_preserves_override_compatibility(self) -> None: + from winml.modelkit.config.build import _merge_export_config + + base = WinMLExportConfig() + override = WinMLExportConfig( + compatibility=ExportCompatibilityConfig(transformers_attention="eager") + ) + + merged = _merge_export_config(base, override) + + assert merged.compatibility.transformers_attention == "eager" + + class TestGetIoSpecsFromConfig: """Unit tests for resolve_io_specs function.""" diff --git a/tests/unit/export/test_config_validation.py b/tests/unit/export/test_config_validation.py index 00d825971..d1f1ed621 100644 --- a/tests/unit/export/test_config_validation.py +++ b/tests/unit/export/test_config_validation.py @@ -19,6 +19,7 @@ OutputTensorSpec, WinMLExportConfig, ) +from winml.modelkit.export.policy import ExportCompatibilityConfig # ============================================================================= @@ -395,6 +396,37 @@ def test_full_roundtrip(self): assert restored.dynamic_axes == {"pixel_values": {1: "channels"}} + +class TestExportCompatibilitySerialization: + def test_compatibility_round_trips_when_present(self) -> None: + cfg = WinMLExportConfig( + compatibility=ExportCompatibilityConfig(transformers_attention="eager") + ) + + data = cfg.to_dict() + round_tripped = WinMLExportConfig.from_dict(data) + + assert data["compatibility"] == {"transformers_attention": "eager"} + assert round_tripped.compatibility.transformers_attention == "eager" + + def test_unresolved_empty_compatibility_is_omitted_from_export_dict(self) -> None: + cfg = WinMLExportConfig() + + assert "compatibility" not in cfg.to_dict() + + def test_empty_compatibility_is_omitted_from_export_dict(self) -> None: + cfg = WinMLExportConfig.from_dict({"compatibility": {}}) + + data = cfg.to_dict() + round_tripped = WinMLExportConfig.from_dict(data) + + assert "compatibility" not in data + assert round_tripped.compatibility.transformers_attention is None + + def test_invalid_compatibility_value_raises(self) -> None: + with pytest.raises(ValueError, match="transformers_attention"): + WinMLExportConfig.from_dict({"compatibility": {"transformers_attention": "sdpa"}}) + def test_from_dict_ignores_unknown_fields(self): data = {"opset_version": 17, "batch_size": 1, "unknown_field": "ignored"} cfg = WinMLExportConfig.from_dict(data) diff --git a/tests/unit/export/test_export_policy.py b/tests/unit/export/test_export_policy.py new file mode 100644 index 000000000..870288169 --- /dev/null +++ b/tests/unit/export/test_export_policy.py @@ -0,0 +1,119 @@ +# ------------------------------------------------------------------------- +# Copyright (c) Microsoft Corporation. All rights reserved. +# Licensed under the MIT License. +# -------------------------------------------------------------------------- + +from __future__ import annotations + +import pytest + +from winml.modelkit.export.policy import ( + ExportCompatibilityConfig, + ExportCompatibilityRule, + ExportPolicyTarget, + _rule_from_dict, + export_policy_targets_for_request, + resolve_export_compatibility, +) + + +def test_qnn_gpu_target_requires_eager_transformers_attention() -> None: + cfg = resolve_export_compatibility([ExportPolicyTarget(ep="qnn", device="gpu")]) + + assert cfg.transformers_attention == "eager" + + +def test_default_rules_load_from_json_catalog() -> None: + from winml.modelkit.export import policy + + rules = policy.load_export_compatibility_rules() + + assert rules == ( + ExportCompatibilityRule( + ep=None, + device=None, + compatibility=ExportCompatibilityConfig(transformers_attention="eager"), + reason="Transformers SDPA-exported attention guard paths are not broadly portable.", + ), + ) + + +def test_global_rule_forces_transformers_attention_for_non_qnn_target() -> None: + cfg = resolve_export_compatibility( + [ExportPolicyTarget(ep="DmlExecutionProvider", device="gpu")] + ) + + assert cfg.transformers_attention == "eager" + + +def test_no_targets_uses_supported_catalog_and_includes_qnn_requirement() -> None: + cfg = resolve_export_compatibility() + + assert cfg.transformers_attention == "eager" + + +def test_export_policy_targets_for_request_keeps_portable_default_when_not_explicit() -> None: + targets = export_policy_targets_for_request( + ep="QNNExecutionProvider", + device="gpu", + target_was_explicit=False, + ) + + assert targets is None + + +def test_export_policy_targets_for_request_resolves_explicit_alias() -> None: + targets = export_policy_targets_for_request( + ep="qnn", + device="gpu", + target_was_explicit=True, + ) + + assert targets == (ExportPolicyTarget(ep="QNNExecutionProvider", device="gpu"),) + + +def test_conflicting_rules_raise_clear_error() -> None: + rules = ( + ExportCompatibilityRule( + ep="QNNExecutionProvider", + device="gpu", + compatibility=ExportCompatibilityConfig(transformers_attention="eager"), + reason="first rule", + ), + ExportCompatibilityRule( + ep="QNNExecutionProvider", + device="gpu", + compatibility=ExportCompatibilityConfig(transformers_attention="sdpa"), # type: ignore[arg-type] + reason="second rule", + ), + ) + + with pytest.raises(ValueError, match="Conflicting export compatibility"): + resolve_export_compatibility( + [ExportPolicyTarget(ep="qnn", device="gpu")], + rules=rules, + ) + + +def test_json_rule_rejects_unknown_ep() -> None: + with pytest.raises(ValueError, match=r"match\.ep"): + _rule_from_dict( + { + "match": {"ep": "TypoExecutionProvider", "device": None}, + "export": {"transformers_attention": "eager"}, + "reason": "typo should not create a dead rule", + }, + index=0, + ) + + +def test_json_rule_rejects_unsupported_device_for_ep() -> None: + with pytest.raises(ValueError, match="does not support device"): + _rule_from_dict( + { + "match": {"ep": "CPUExecutionProvider", "device": "gpu"}, + "export": {"transformers_attention": "eager"}, + "reason": "device mismatch should not create a dead rule", + }, + index=0, + ) diff --git a/tests/unit/export/test_htp_exporter_attention_compat.py b/tests/unit/export/test_htp_exporter_attention_compat.py new file mode 100644 index 000000000..a2b20d14f --- /dev/null +++ b/tests/unit/export/test_htp_exporter_attention_compat.py @@ -0,0 +1,299 @@ +# ------------------------------------------------------------------------- +# Copyright (c) Microsoft Corporation. All rights reserved. +# Licensed under the MIT License. +# -------------------------------------------------------------------------- +"""Tests for export-time attention compatibility handling.""" + +from __future__ import annotations + +from typing import TYPE_CHECKING + +import torch +import torch.nn as nn + +from winml.modelkit.export import InputTensorSpec, OutputTensorSpec, WinMLExportConfig +from winml.modelkit.export.htp import HTPExporter +from winml.modelkit.export.policy import ExportCompatibilityConfig + + +if TYPE_CHECKING: + from pathlib import Path + + import pytest + + +class _AttentionConfig: + """Minimal HF-style config with an attention implementation knob.""" + + model_type = "fake" + + def __init__(self, implementation: str = "sdpa") -> None: + self._attn_implementation = implementation + + +class _CascadingAttentionConfig: + """HF-style config whose setter cascades attention into child configs.""" + + model_type = "fake" + + def __init__( + self, + implementation: str, + *, + sub_configs: list[_CascadingAttentionConfig] + | dict[str, _CascadingAttentionConfig] + | None = None, + ) -> None: + self._implementation = implementation + self.sub_configs = sub_configs or [] + + @property + def _attn_implementation(self) -> str: + return self._implementation + + @_attn_implementation.setter + def _attn_implementation(self, value: str) -> None: + self._implementation = value + sub_configs = ( + self.sub_configs.values() if isinstance(self.sub_configs, dict) else self.sub_configs + ) + for config in sub_configs: + config._attn_implementation = value + + +class _NestedAttentionModel(nn.Module): + """Model with root and child configs to mirror HF module trees.""" + + def __init__(self) -> None: + super().__init__() + self.config = _AttentionConfig() + self.proj = nn.Linear(2, 2) + self.proj.config = _AttentionConfig() + + def forward(self, x: torch.Tensor) -> torch.Tensor: + return self.proj(x) + + +class _CascadingAttentionModel(nn.Module): + """Model shaped like HF composite configs where the root setter recurses.""" + + def __init__(self) -> None: + super().__init__() + child_config = _CascadingAttentionConfig("flash_attention_2") + self.config = _CascadingAttentionConfig("sdpa", sub_configs=[child_config]) + self.proj = nn.Linear(2, 2) + self.proj.config = child_config + + def forward(self, x: torch.Tensor) -> torch.Tensor: + return self.proj(x) + + +class _CascadingEagerChildAttentionModel(nn.Module): + """Composite config with an eager child that still needs restore snapshotting.""" + + def __init__(self) -> None: + super().__init__() + child_config = _CascadingAttentionConfig("eager") + self.config = _CascadingAttentionConfig("sdpa", sub_configs=[child_config]) + self.proj = nn.Linear(2, 2) + self.proj.config = child_config + + def forward(self, x: torch.Tensor) -> torch.Tensor: + return self.proj(x) + + +class _CascadingDetachedEagerChildAttentionModel(nn.Module): + """Composite config whose child is reachable only through parent.sub_configs.""" + + def __init__(self) -> None: + super().__init__() + self.child_config = _CascadingAttentionConfig("eager") + self.config = _CascadingAttentionConfig( + "sdpa", sub_configs={"child_config": self.child_config} + ) + self.proj = nn.Linear(2, 2) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + return self.proj(x) + + +class _ChildBeforeParentAttentionModel(nn.Module): + """Model where module traversal sees a child config before its cascading parent.""" + + def __init__(self) -> None: + super().__init__() + self.child_config = _CascadingAttentionConfig("eager") + self.early = nn.Linear(2, 2) + self.early.config = self.child_config + self.parent_config = _CascadingAttentionConfig("sdpa", sub_configs=[self.child_config]) + self.late = nn.Linear(2, 2) + self.late.config = self.parent_config + + def forward(self, x: torch.Tensor) -> torch.Tensor: + return self.late(self.early(x)) + + +def _export_config(*, eager_attention: bool) -> WinMLExportConfig: + return WinMLExportConfig( + input_tensors=[InputTensorSpec(name="x", dtype="float32", shape=(1, 2))], + output_tensors=[OutputTensorSpec(name="y")], + compatibility=ExportCompatibilityConfig( + transformers_attention="eager" if eager_attention else None + ), + ) + + +def test_htp_exporter_uses_eager_attention_when_policy_requests_it( + monkeypatch: pytest.MonkeyPatch, + tmp_path: Path, +) -> None: + model = _NestedAttentionModel() + captured: dict[str, str] = {} + + def fake_export(*args: object, **kwargs: object) -> None: + captured["root"] = model.config._attn_implementation + captured["child"] = model.proj.config._attn_implementation + + monkeypatch.setattr(torch.onnx, "export", fake_export) + + HTPExporter()._convert_model_to_onnx( + model, + str(tmp_path / "model.onnx"), + {"x": torch.ones(1, 2)}, + _export_config(eager_attention=True), + task=None, + ) + + assert captured == {"root": "eager", "child": "eager"} + assert model.config._attn_implementation == "sdpa" + assert model.proj.config._attn_implementation == "sdpa" + + +def test_htp_exporter_restores_nested_attention_configs_losslessly( + monkeypatch: pytest.MonkeyPatch, + tmp_path: Path, +) -> None: + model = _CascadingAttentionModel() + captured: dict[str, str] = {} + + def fake_export(*args: object, **kwargs: object) -> None: + captured["root"] = model.config._attn_implementation + captured["child"] = model.proj.config._attn_implementation + + monkeypatch.setattr(torch.onnx, "export", fake_export) + + HTPExporter()._convert_model_to_onnx( + model, + str(tmp_path / "model.onnx"), + {"x": torch.ones(1, 2)}, + _export_config(eager_attention=True), + task=None, + ) + + assert captured == {"root": "eager", "child": "eager"} + assert model.config._attn_implementation == "sdpa" + assert model.proj.config._attn_implementation == "flash_attention_2" + + +def test_htp_exporter_restores_eager_child_after_parent_restore_cascades( + monkeypatch: pytest.MonkeyPatch, + tmp_path: Path, +) -> None: + model = _CascadingEagerChildAttentionModel() + captured: dict[str, str] = {} + + def fake_export(*args: object, **kwargs: object) -> None: + captured["root"] = model.config._attn_implementation + captured["child"] = model.proj.config._attn_implementation + + monkeypatch.setattr(torch.onnx, "export", fake_export) + + HTPExporter()._convert_model_to_onnx( + model, + str(tmp_path / "model.onnx"), + {"x": torch.ones(1, 2)}, + _export_config(eager_attention=True), + task=None, + ) + + assert captured == {"root": "eager", "child": "eager"} + assert model.config._attn_implementation == "sdpa" + assert model.proj.config._attn_implementation == "eager" + + +def test_htp_exporter_restores_sub_config_not_attached_to_module( + monkeypatch: pytest.MonkeyPatch, + tmp_path: Path, +) -> None: + model = _CascadingDetachedEagerChildAttentionModel() + captured: dict[str, str] = {} + + def fake_export(*args: object, **kwargs: object) -> None: + captured["root"] = model.config._attn_implementation + captured["child"] = model.child_config._attn_implementation + + monkeypatch.setattr(torch.onnx, "export", fake_export) + + HTPExporter()._convert_model_to_onnx( + model, + str(tmp_path / "model.onnx"), + {"x": torch.ones(1, 2)}, + _export_config(eager_attention=True), + task=None, + ) + + assert captured == {"root": "eager", "child": "eager"} + assert model.config._attn_implementation == "sdpa" + assert model.child_config._attn_implementation == "eager" + + +def test_htp_exporter_restores_child_discovered_before_parent( + monkeypatch: pytest.MonkeyPatch, + tmp_path: Path, +) -> None: + model = _ChildBeforeParentAttentionModel() + captured: dict[str, str] = {} + + def fake_export(*args: object, **kwargs: object) -> None: + captured["parent"] = model.parent_config._attn_implementation + captured["child"] = model.child_config._attn_implementation + + monkeypatch.setattr(torch.onnx, "export", fake_export) + + HTPExporter()._convert_model_to_onnx( + model, + str(tmp_path / "model.onnx"), + {"x": torch.ones(1, 2)}, + _export_config(eager_attention=True), + task=None, + ) + + assert captured == {"parent": "eager", "child": "eager"} + assert model.parent_config._attn_implementation == "sdpa" + assert model.child_config._attn_implementation == "eager" + + +def test_htp_exporter_leaves_attention_unchanged_without_policy( + monkeypatch: pytest.MonkeyPatch, + tmp_path: Path, +) -> None: + model = _NestedAttentionModel() + captured: dict[str, str] = {} + + def fake_export(*args: object, **kwargs: object) -> None: + captured["root"] = model.config._attn_implementation + captured["child"] = model.proj.config._attn_implementation + + monkeypatch.setattr(torch.onnx, "export", fake_export) + + HTPExporter()._convert_model_to_onnx( + model, + str(tmp_path / "model.onnx"), + {"x": torch.ones(1, 2)}, + _export_config(eager_attention=False), + task=None, + ) + + assert captured == {"root": "sdpa", "child": "sdpa"} + assert model.config._attn_implementation == "sdpa" + assert model.proj.config._attn_implementation == "sdpa" diff --git a/tests/unit/models/auto/test_auto_onnx.py b/tests/unit/models/auto/test_auto_onnx.py index 0e0b328e8..8959d2609 100644 --- a/tests/unit/models/auto/test_auto_onnx.py +++ b/tests/unit/models/auto/test_auto_onnx.py @@ -393,6 +393,91 @@ def test_delegates_uppercase_local_onnx_to_from_onnx( assert from_onnx.call_args.kwargs["onnx_path"] == local_onnx +class TestFromPretrainedBuildConfigTarget: + """Target request forwarding for HF from_pretrained builds.""" + + def test_config_uses_runtime_target_and_requested_export_policy(self, tmp_path: Path) -> None: + """Runtime target drives quant/compile while request target drives export policy.""" + from winml.modelkit.config import WinMLBuildConfig + from winml.modelkit.loader import WinMLLoaderConfig + + ep_device = MagicMock() + ep_device.device.ep_name = "QNNExecutionProvider" + ep_device.device.device_type = "GPU" + build_config = WinMLBuildConfig( + loader=WinMLLoaderConfig(task="image-classification", model_type="resnet"), + compile=None, + ) + hf_config = MagicMock() + hf_config.model_type = "resnet" + build_result = _make_build_result(tmp_path) + + with ( + patch( + "winml.modelkit.config.generate_hf_build_config", + return_value=build_config, + ) as mock_gen, + patch("winml.modelkit.loader.load_hf_config", return_value=hf_config), + patch("winml.modelkit.build.build_hf_model", return_value=build_result), + patch( + "winml.modelkit.models.auto.get_winml_class", + return_value=lambda **_: MagicMock(), + ), + ): + WinMLAutoModel.from_pretrained( + "fake/model", + ep_device=ep_device, + device="auto", + ep="qnn", + task="image-classification", + ) + + assert mock_gen.call_args.kwargs["device"] == "gpu" + assert mock_gen.call_args.kwargs["ep"] == "qnn" + assert mock_gen.call_args.kwargs["export_policy_target"] == ("auto", "qnn") + + def test_bare_default_request_stays_auto_after_runtime_resolution(self, tmp_path: Path) -> None: + """Default export policy stays portable while quant/compile use runtime target.""" + from winml.modelkit.config import WinMLBuildConfig + from winml.modelkit.loader import WinMLLoaderConfig + from winml.modelkit.session import EPDeviceTarget + + ep_device = MagicMock() + ep_device.device.ep_name = "DmlExecutionProvider" + ep_device.device.device_type = "GPU" + build_config = WinMLBuildConfig( + loader=WinMLLoaderConfig(task="image-classification", model_type="resnet"), + compile=None, + ) + hf_config = MagicMock() + hf_config.model_type = "resnet" + build_result = _make_build_result(tmp_path) + + with ( + patch( + "winml.modelkit.session.resolve_device", + return_value=EPDeviceTarget(ep="DmlExecutionProvider", device="gpu"), + ), + patch("winml.modelkit.session.WinMLEPRegistry.instance") as mock_registry, + patch( + "winml.modelkit.config.generate_hf_build_config", + return_value=build_config, + ) as mock_gen, + patch("winml.modelkit.loader.load_hf_config", return_value=hf_config), + patch("winml.modelkit.build.build_hf_model", return_value=build_result), + patch( + "winml.modelkit.models.auto.get_winml_class", + return_value=lambda **_: MagicMock(), + ), + ): + mock_registry.return_value.auto_device.return_value = ep_device + WinMLAutoModel.from_pretrained("fake/model", task="image-classification") + + assert mock_gen.call_args.kwargs["device"] == "gpu" + assert mock_gen.call_args.kwargs["ep"] == "dml" + assert mock_gen.call_args.kwargs["export_policy_target"] == ("auto", None) + + # ============================================================================= # from_onnx cache dir and cache_key tests # ============================================================================= diff --git a/tests/unit/models/auto/test_composite_model_type_routing.py b/tests/unit/models/auto/test_composite_model_type_routing.py index 6f86d2a2b..a3af47d59 100644 --- a/tests/unit/models/auto/test_composite_model_type_routing.py +++ b/tests/unit/models/auto/test_composite_model_type_routing.py @@ -76,7 +76,9 @@ class _Cfg: lambda *args, **kwargs: _Cfg(), ) - ep_device = SimpleNamespace(device=SimpleNamespace(device_type="CPU")) + ep_device = SimpleNamespace( + device=SimpleNamespace(device_type="CPU", ep_name="CPUExecutionProvider") + ) result = WinMLAutoModel.from_pretrained( "dummy/model", task=_SHARED_TASK, ep_device=ep_device ) @@ -100,7 +102,9 @@ def test_explicit_model_type_routes_without_native_config_probe( ) ), ) - ep_device = SimpleNamespace(device=SimpleNamespace(device_type="CPU")) + ep_device = SimpleNamespace( + device=SimpleNamespace(device_type="CPU", ep_name="CPUExecutionProvider") + ) result = WinMLAutoModel.from_pretrained( "dummy/model", diff --git a/tests/unit/models/winml/test_composite_from_pretrained.py b/tests/unit/models/winml/test_composite_from_pretrained.py index ad8a3bbd9..1f12839ba 100644 --- a/tests/unit/models/winml/test_composite_from_pretrained.py +++ b/tests/unit/models/winml/test_composite_from_pretrained.py @@ -151,6 +151,30 @@ def resolve_target(target: object) -> object: assert targets[0].ep == "qnn" +def test_sub_models_receive_original_request_target() -> None: + """Composite sub-model builds must not infer export policy from the runtime handle.""" + hf_cfg = SimpleNamespace(model_type="_test_model_type") + fake_ep_device = _fake_ep_device() + + with ( + patch("winml.modelkit.loader.load_hf_config", return_value=hf_cfg), + patch( + "winml.modelkit.models.auto.WinMLAutoModel.from_pretrained", + return_value=MagicMock(), + ) as mock_from_pretrained, + ): + _StubComposite.from_pretrained( + "hf/id", + task="_test_task", + device="auto", + ep_device=fake_ep_device, + ) + + call_kwargs = mock_from_pretrained.call_args.kwargs + assert call_kwargs["device"] == "auto" + assert call_kwargs["ep"] is None + + def test_device_is_forwarded_from_from_onnx() -> None: """The resolved component target must also reach the composite wrapper.""" fake_ep_device = _fake_ep_device()