diff --git a/src/winml/modelkit/commands/optimize.py b/src/winml/modelkit/commands/optimize.py index f86092d0e..889fb7003 100644 --- a/src/winml/modelkit/commands/optimize.py +++ b/src/winml/modelkit/commands/optimize.py @@ -229,7 +229,26 @@ def _render_check_optim(console: Console, findings: list[Any], verbose: bool) -> ) -def _run_check_optim(model: Path, all_caps: dict[str, Any], verbose: bool) -> None: +def _resolve_optimization_target(ep: str | None, device: str | None) -> tuple[Any, Any]: + """Resolve an explicit optimization EP/device request.""" + from ..session import ( + EPDeviceTarget, + WinMLEPRegistry, + resolve_device, + ) + + target = resolve_device(EPDeviceTarget(ep=ep or "auto", device=(device or "auto").lower())) + return target, WinMLEPRegistry.instance().auto_device(target) + + +def _run_check_optim( + model: Path, + all_caps: dict[str, Any], + verbose: bool, + *, + target: Any | None = None, + ep_device: Any | None = None, +) -> None: """Probe which optimizations apply to the model and print a report. No output file is written. Every boolean capability that is off by default @@ -239,6 +258,8 @@ def _run_check_optim(model: Path, all_caps: dict[str, Any], verbose: bool) -> No model: Path to the input ONNX model. all_caps: The full capability registry. verbose: Whether to show every affected node/constant. + target: Resolved EP/device target, if explicitly requested. + ep_device: Resolved runtime EP device used by each optimization probe. """ from ..optim import BoolCapability, analyze_model @@ -247,6 +268,8 @@ def _run_check_optim(model: Path, all_caps: dict[str, Any], verbose: bool) -> No ) console.print(f"[bold blue]Input:[/bold blue] {model}") + if target is not None: + console.print(f"[bold blue]Target:[/bold blue] {target.ep} on {target.device.upper()}") console.print( "[dim]--check-optim — analyzing applicable optimizations (no output written).[/dim]" ) @@ -268,6 +291,7 @@ def _run_check_optim(model: Path, all_caps: dict[str, Any], verbose: bool) -> No findings = analyze_model( onnx_model, all_caps, + ep_device=ep_device, on_probe_start=lambda name: progress.update(task_id, description=name), on_probe_complete=lambda _: progress.advance(task_id), ) @@ -303,6 +327,19 @@ def _run_check_optim(model: Path, all_caps: dict[str, Any], verbose: bool) -> No ) @cli_utils.output_option("Output path (default: {input}_opt.onnx)") @cli_utils.overwrite_option() +@cli_utils.ep_option( + required=False, + default=None, + include_auto=True, + include_cuda=True, + optional_message="If omitted with --device, selects a compatible EP automatically.", +) +@cli_utils.device_option( + required=False, + default=None, + include_auto=True, + optional_message="If both --ep and --device are omitted, preserves CPU optimization.", +) @click.option( "--config", "-c", @@ -322,6 +359,8 @@ def optimize( model: Path | None, output: Path | None, overwrite: bool, + ep: str | None, + device: str | None, config: Path | None, verbose: int, quiet: bool, @@ -476,9 +515,37 @@ def optimize( verbose, quiet = cli_utils.resolve_verbosity(ctx, verbose, quiet) configure_logging(verbosity=verbose, quiet=quiet) + target = None + ep_device = None + if ep is not None or device is not None: + from ..session import ( + DeviceNotFound, + UnknownListingPick, + WinMLEPNotDiscovered, + WinMLEPRegistrationFailed, + ) + + try: + target, ep_device = _resolve_optimization_target(ep, device) + except ( + DeviceNotFound, + RuntimeError, + UnknownListingPick, + ValueError, + WinMLEPNotDiscovered, + WinMLEPRegistrationFailed, + ) as e: + raise click.UsageError(f"Could not resolve optimization target: {e}") from e + # Handle --check-optim: report which optimizations apply, write nothing. if check_optim: - _run_check_optim(model, all_caps, bool(verbose)) + _run_check_optim( + model, + all_caps, + bool(verbose), + target=target, + ep_device=ep_device, + ) return # Import optimizer @@ -542,6 +609,14 @@ def optimize( if verbose: optimizer_kwargs["verbose"] = True + if ep_device is not None: + if target is None: + raise RuntimeError("Resolved optimization EP device has no target metadata.") + optimizer_kwargs["ep_device"] = ep_device + console.print( + f"[bold blue]Target:[/bold blue] {target.ep} on {target.device.upper()}" + ) + try: console.print("\n[bold]Loading model...[/bold]") onnx_model = load_onnx(model) @@ -556,10 +631,22 @@ def optimize( # Report results optimized_nodes = len(optimized_model.graph.node) - reduction = (1 - optimized_nodes / original_nodes) * 100 if original_nodes else 0 + node_change = ( + abs(optimized_nodes - original_nodes) / original_nodes * 100 if original_nodes else 0 + ) console.print(f"\n[bold green]Success![/bold green] Model optimized: {output}") - node_info = f"Nodes: {original_nodes} -> {optimized_nodes} ({reduction:.1f}% reduction)" + if optimized_nodes < original_nodes: + change_label = "reduction" + elif optimized_nodes > original_nodes: + change_label = "increase" + else: + change_label = None + node_info = f"Nodes: {original_nodes} -> {optimized_nodes}" + if change_label is None: + node_info += " (no change)" + else: + node_info += f" ({node_change:.1f}% {change_label})" console.print(f"[dim]{node_info}[/dim]") except Exception as e: diff --git a/src/winml/modelkit/optim/analysis.py b/src/winml/modelkit/optim/analysis.py index cbea1b3a1..f022f1716 100644 --- a/src/winml/modelkit/optim/analysis.py +++ b/src/winml/modelkit/optim/analysis.py @@ -283,8 +283,10 @@ def _run_pipe(pipe: Any, model: ModelProto, config: Any) -> ModelProto: def _iter_findings( model: ModelProto, capabilities: dict[str, CapabilityDef], + *, on_probe_start: Callable[[str], None] | None = None, on_probe_complete: Callable[[str], None] | None = None, + **optimizer_kwargs: Any, ) -> Iterator[tuple[CapabilityFinding, ModelProto]]: """Yield ``(finding, produced_model)`` for every applicable optimization. @@ -319,6 +321,7 @@ def complete_probe(cap_name: str) -> None: # Baseline kwargs = every capability at its default value. default_kwargs = {cap.python_name: cap.default for cap in capabilities.values()} + default_kwargs.update(optimizer_kwargs) kebab_defaults = {name: cap.default for name, cap in capabilities.items()} # Mandatory pre-stage — mirrors Optimizer.optimize(). The clone is required: @@ -358,6 +361,22 @@ def complete_probe(cap_name: str) -> None: if on_probe_start is not None: on_probe_start(cap_name) try: + ep_device = optimizer_kwargs.get("ep_device") + if ep_device is not None and cap.ep_constraint is not None: + from ..utils.constants import normalize_ep_name + + target_ep = normalize_ep_name(ep_device.device.ep_name) + if not any( + normalize_ep_name(name) == target_ep for name in cap.ep_constraint + ): + logger.debug( + "Skipping capability '%s': target EP %s is not in %s", + cap.name, + target_ep, + cap.ep_constraint, + ) + continue + # Enable only this capability (plus its dependencies) on top of # the all-defaults configuration. kebab = dict(kebab_defaults) @@ -368,6 +387,7 @@ def complete_probe(cap_name: str) -> None: for name, value in kebab.items() if name in capabilities } + probe_kwargs.update(optimizer_kwargs) probe_config = pipe.build_config(**probe_kwargs) should_process = getattr(pipe, "should_process", None) @@ -428,6 +448,7 @@ def analyze_model( *, on_probe_start: Callable[[str], None] | None = None, on_probe_complete: Callable[[str], None] | None = None, + **optimizer_kwargs: Any, ) -> list[CapabilityFinding]: """Probe every applicable optimization capability against ``model``. @@ -440,6 +461,8 @@ def analyze_model( model: The input ONNX model (never modified). capabilities: The full capability registry (kebab-case keyed), e.g. from ``optim.pipes.get_all_capabilities()``. + **optimizer_kwargs: Pipeline context forwarded to every pipe configuration, + such as a resolved ``ep_device``. on_probe_start: Optional callback invoked with the capability name before each probe begins. on_probe_complete: Optional callback invoked with the capability name @@ -456,6 +479,7 @@ def analyze_model( capabilities, on_probe_start=on_probe_start, on_probe_complete=on_probe_complete, + **optimizer_kwargs, ) ] @@ -463,6 +487,7 @@ def analyze_model( def iter_optimization_outputs( model: ModelProto, capabilities: dict[str, CapabilityDef], + **optimizer_kwargs: Any, ) -> Iterator[tuple[CapabilityFinding, ModelProto]]: """Yield each applicable optimization together with the model it produces. @@ -476,10 +501,11 @@ def iter_optimization_outputs( Args: model: The input ONNX model (never modified). capabilities: The full capability registry (kebab-case keyed). + **optimizer_kwargs: Pipeline context forwarded to every pipe configuration. Yields: ``(finding, produced_model)`` pairs in pipeline order, one per applicable optimization. The pairs are produced lazily; materialize the iterator if the produced models must outlive iteration. """ - yield from _iter_findings(model, capabilities) + yield from _iter_findings(model, capabilities, **optimizer_kwargs) diff --git a/src/winml/modelkit/optim/pipes/graph.py b/src/winml/modelkit/optim/pipes/graph.py index 853cdb9ee..c1c82e9e1 100644 --- a/src/winml/modelkit/optim/pipes/graph.py +++ b/src/winml/modelkit/optim/pipes/graph.py @@ -26,6 +26,8 @@ if TYPE_CHECKING: import onnx + from ...session import WinMLEPDevice + # Import all capability modules to build capabilities dict from ..capabilities import ( activation, @@ -146,12 +148,15 @@ def __init__( self, enabled: list[str] | None = None, verbose: bool = False, + ep_device: WinMLEPDevice | None = None, ) -> None: """Initialize with all advanced optimizers disabled, enable specified ones. Args: enabled: List of python_names to enable (e.g., ["gelu_fusion"]) verbose: Enable verbose logging + ep_device: Resolved EP/device target for provider-aware optimization. + If omitted, optimization uses CPUExecutionProvider. Note: Only default=False capabilities are managed here. Basic ORT @@ -159,6 +164,7 @@ def __init__( """ self.optimization_level = 2 # Finalized at Level 2 self.verbose = verbose + self.ep_device = ep_device # Special flag for GeluApproximation (requires separate session config) # ORT docs: "GeluApproximation has side effects which may change results. @@ -356,7 +362,11 @@ def build_config(cls, **kwargs: Any) -> ORTGraphPipeConfig: and cap.python_name not in explicitly_disabled ] - config = ORTGraphPipeConfig(enabled=enabled, verbose=verbose) + config = ORTGraphPipeConfig( + enabled=enabled, + verbose=verbose, + ep_device=kwargs.get("ep_device"), + ) # Explicitly disable capabilities that user set to False # This handles default=True caps like constant_folding @@ -456,7 +466,12 @@ def _log_process_verbose( " graph_optimization_level: %d (ORT_ENABLE_EXTENDED)", config.optimization_level ) logger.debug(" optimized_model_filepath: %s", output_file) - logger.debug(" providers: ['CPUExecutionProvider']") + provider = ( + config.ep_device.device.ep_name + if config.ep_device is not None + else "CPUExecutionProvider" + ) + logger.debug(" provider: %s", provider) # Session config entries logger.debug("[Session Config Entries]") @@ -562,15 +577,38 @@ def process(self, model: onnx.ModelProto, config: ORTGraphPipeConfig) -> onnx.Mo "1", ) + if config.ep_device is not None: + from ...session import lookup_device_spec + + spec = lookup_device_spec( + config.ep_device.device.ep_name, + config.ep_device.device.device_type.lower(), + ) + if spec is None: + logger.debug( + "No device specification found for %s on %s; " + "using empty provider options", + config.ep_device.device.ep_name, + config.ep_device.device.device_type, + ) + provider_options = dict(spec.default_provider_options) if spec else {} + sess_opts.add_provider_for_devices( + [config.ep_device.device.ort_handle], + provider_options, + ) + # Verbose output for process if config.verbose: self._log_process_verbose(config, model, input_file, output_file, disable_list) # Create session to trigger optimization try: - _ = ort.InferenceSession( - str(input_file), sess_opts, providers=["CPUExecutionProvider"] - ) + if config.ep_device is None: + session = ort.InferenceSession( + str(input_file), sess_opts, providers=["CPUExecutionProvider"] + ) + else: + session = ort.InferenceSession(str(input_file), sess_opts) except Exception as e: raise OptimizationError( f"ONNX Runtime optimization failed: {e}", @@ -582,6 +620,22 @@ def process(self, model: onnx.ModelProto, config: ORTGraphPipeConfig) -> onnx.Mo cause=e, ) from e + if config.ep_device is not None: + from ...utils.constants import normalize_ep_name + + expected_provider = normalize_ep_name(config.ep_device.device.ep_name) + active_providers = session.get_providers() + logger.debug("ORT optimization session providers: %s", active_providers) + if not any( + normalize_ep_name(provider) == expected_provider + for provider in active_providers + ): + raise OptimizationError( + f"Requested provider {expected_provider} was not activated; " + f"active providers: {active_providers}", + pipe_name=self.name, + ) + # Load and return optimized model try: return load_onnx(output_file, validate=False) diff --git a/src/winml/modelkit/utils/cli.py b/src/winml/modelkit/utils/cli.py index 46da14c8a..d1a12049c 100644 --- a/src/winml/modelkit/utils/cli.py +++ b/src/winml/modelkit/utils/cli.py @@ -415,6 +415,7 @@ def ep_option( default: str | None = None, include_auto: bool = False, include_all: bool = False, + include_cuda: bool = False, ) -> Callable[[F], F]: """Add --ep (execution provider) option to a Click command. @@ -428,6 +429,8 @@ def ep_option( (default: False). include_all: Whether to include "all" as a valid choice (default: False). + include_cuda: Whether to include CUDA aliases and the full provider name + (default: False). Returns: Decorator function @@ -440,7 +443,11 @@ def ep_option( if optional_message: help_text = f"{help_text}. {optional_message}" - ep_choices = [name for name in ALL_EP_NAMES if name not in ("cuda", "CUDAExecutionProvider")] + ep_choices = [ + name + for name in ALL_EP_NAMES + if include_cuda or name not in ("cuda", "CUDAExecutionProvider") + ] choices = ["auto", *ep_choices] if include_auto else ep_choices choices = ["all", *choices] if include_all else choices diff --git a/tests/unit/commands/test_optimize_cli.py b/tests/unit/commands/test_optimize_cli.py index 98dda97ca..314c91b74 100644 --- a/tests/unit/commands/test_optimize_cli.py +++ b/tests/unit/commands/test_optimize_cli.py @@ -73,7 +73,19 @@ def test_help_exits_cleanly(self, runner: CliRunner) -> None: def test_help_shows_required_flags(self, runner: CliRunner) -> None: result = runner.invoke(optimize, ["--help"]) assert result.exit_code == 0 - for flag in ("--model", "-m", "--output", "-o", "--config", "-c", "--verbose", "-v"): + for flag in ( + "--model", + "-m", + "--output", + "-o", + "--ep", + "--device", + "-d", + "--config", + "-c", + "--verbose", + "-v", + ): assert flag in result.output, f"Missing flag {flag} in help" def test_model_required_without_list_flags(self, runner: CliRunner) -> None: @@ -189,6 +201,72 @@ def test_node_reduction_reported(self, runner: CliRunner, tmp_path: Path) -> Non assert result.exit_code == 0, result.output assert "10" in result.output assert "8" in result.output + assert "20.0% reduction" in result.output + + def test_node_increase_reported(self, runner: CliRunner, tmp_path: Path) -> None: + model_file = tmp_path / "model.onnx" + model_file.touch() + + original = _make_mock_model(num_nodes=10) + optimized = _make_mock_model(num_nodes=12) + with ( + patch(_LOAD_ONNX, return_value=original), + patch(_SAVE_ONNX), + patch(_OPTIMIZER) as mock_opt_cls, + ): + mock_opt_cls.return_value.optimize.return_value = optimized + result = runner.invoke(optimize, ["-m", str(model_file)]) + + assert result.exit_code == 0, result.output + assert "20.0% increase" in result.output + + def test_unchanged_node_count_reported(self, runner: CliRunner, tmp_path: Path) -> None: + model_file = tmp_path / "model.onnx" + model_file.touch() + model = _make_mock_model(num_nodes=10) + + with ( + patch(_LOAD_ONNX, return_value=model), + patch(_SAVE_ONNX), + patch(_OPTIMIZER) as mock_opt_cls, + ): + mock_opt_cls.return_value.optimize.return_value = model + result = runner.invoke(optimize, ["-m", str(model_file)]) + + assert result.exit_code == 0, result.output + assert "Nodes: 10 -> 10 (no change)" in result.output + + def test_device_target_forwarded_to_optimizer( + self, runner: CliRunner, tmp_path: Path + ) -> None: + model_file = tmp_path / "model.onnx" + model_file.touch() + mock_model = _make_mock_model() + resolved_target = MagicMock(ep="DmlExecutionProvider", device="gpu") + resolved_ep_device = MagicMock() + + with ( + patch(_LOAD_ONNX, return_value=mock_model), + patch(_SAVE_ONNX), + patch(_OPTIMIZER) as mock_opt_cls, + patch( + "winml.modelkit.session.resolve_device", + return_value=resolved_target, + ) as mock_resolve, + patch("winml.modelkit.session.WinMLEPRegistry") as mock_registry_cls, + ): + mock_registry_cls.instance.return_value.auto_device.return_value = resolved_ep_device + mock_opt_cls.return_value.optimize.return_value = mock_model + result = runner.invoke(optimize, ["-m", str(model_file), "--device", "gpu"]) + + assert result.exit_code == 0, result.output + requested_target = mock_resolve.call_args.args[0] + assert requested_target.ep == "auto" + assert requested_target.device == "gpu" + assert ( + mock_opt_cls.return_value.optimize.call_args.kwargs["ep_device"] + is resolved_ep_device + ) # ============================================================================= @@ -270,6 +348,49 @@ def test_check_optim_no_findings_message(self, runner: CliRunner, tmp_path: Path assert result.exit_code == 0, result.output assert "No registered optimizations" in result.output + def test_check_optim_forwards_resolved_target( + self, runner: CliRunner, tmp_path: Path + ) -> None: + model_file = tmp_path / "model.onnx" + model_file.touch() + target = MagicMock(ep="DmlExecutionProvider", device="gpu") + ep_device = MagicMock() + + with ( + patch(_LOAD_ONNX, return_value=_make_mock_model()), + patch(_ANALYZE_MODEL, return_value=[]) as mock_analyze, + patch( + "winml.modelkit.commands.optimize._resolve_optimization_target", + return_value=(target, ep_device), + ) as mock_resolve, + ): + result = runner.invoke( + optimize, + ["-m", str(model_file), "--check-optim", "--device", "gpu"], + ) + + assert result.exit_code == 0, result.output + mock_resolve.assert_called_once_with(None, "gpu") + assert mock_analyze.call_args.kwargs["ep_device"] is ep_device + assert "DmlExecutionProvider on GPU" in result.output + + def test_check_optim_accepts_cuda_ep(self, runner: CliRunner, tmp_path: Path) -> None: + model_file = tmp_path / "model.onnx" + model_file.touch() + + with patch( + "winml.modelkit.commands.optimize._resolve_optimization_target", + side_effect=ValueError("CUDA unavailable"), + ) as mock_resolve: + result = runner.invoke( + optimize, + ["-m", str(model_file), "--check-optim", "--ep", "cuda"], + ) + + assert result.exit_code != 0 + mock_resolve.assert_called_once_with("cuda", None) + assert "Invalid value for '--ep'" not in result.output + def test_check_optim_shows_probe_name_and_completed_progress( self, runner: CliRunner, tmp_path: Path ) -> None: @@ -287,10 +408,12 @@ def analyze_with_progress( model: MagicMock, capabilities: dict[str, object], *, + ep_device: object, on_probe_start: object, on_probe_complete: object, ) -> list[object]: del model + assert ep_device is None assert callable(on_probe_start) assert callable(on_probe_complete) for name, cap in capabilities.items(): diff --git a/tests/unit/optim/pipes/test_pipe_config.py b/tests/unit/optim/pipes/test_pipe_config.py index 4e6093daf..8f0be775f 100644 --- a/tests/unit/optim/pipes/test_pipe_config.py +++ b/tests/unit/optim/pipes/test_pipe_config.py @@ -21,20 +21,19 @@ from __future__ import annotations import logging -from typing import TYPE_CHECKING +from unittest.mock import MagicMock, patch + +import pytest from winml.modelkit.optim.pipes import ( GRAPH_CAPABILITIES, + OptimizationError, ORTGraphPipe, ORTGraphPipeConfig, PipeConfig, ) -if TYPE_CHECKING: - import pytest - - # ============================================================================= # TEST CONSTANTS - Capabilities used for testing (all default=False) # ============================================================================= @@ -131,6 +130,70 @@ def test_verbose_parameter(self) -> None: assert config_verbose.verbose is True assert config_quiet.verbose is False + def test_ep_device_parameter(self) -> None: + ep_device = MagicMock() + + config = ORTGraphPipeConfig(ep_device=ep_device) + + assert config.ep_device is ep_device + + def test_build_config_forwards_ep_device(self) -> None: + ep_device = MagicMock() + + config = ORTGraphPipe.build_config(ep_device=ep_device) + + assert config.ep_device is ep_device + + def test_process_binds_resolved_ep_device(self, caplog: pytest.LogCaptureFixture) -> None: + model = MagicMock() + optimized_model = MagicMock() + ep_device = MagicMock() + ep_device.device.ep_name = "DmlExecutionProvider" + ep_device.device.device_type = "GPU" + handle = ep_device.device.ort_handle + session_options = MagicMock() + config = ORTGraphPipeConfig(ep_device=ep_device) + caplog.set_level("DEBUG", logger="winml.modelkit.optim.pipes.graph") + + with ( + patch("onnxruntime.SessionOptions", return_value=session_options), + patch("onnxruntime.InferenceSession") as inference_session, + patch("winml.modelkit.optim.pipes.graph.save_onnx"), + patch( + "winml.modelkit.optim.pipes.graph.load_onnx", + return_value=optimized_model, + ), + patch("winml.modelkit.session.lookup_device_spec", return_value=None), + ): + inference_session.return_value.get_providers.return_value = [ + ep_device.device.ep_name + ] + result = ORTGraphPipe().process(model, config) + + assert result is optimized_model + session_options.add_provider_for_devices.assert_called_once_with([handle], {}) + assert "providers" not in inference_session.call_args.kwargs + assert "using empty provider options" in caplog.text + + def test_process_rejects_provider_fallback(self) -> None: + model = MagicMock() + ep_device = MagicMock() + ep_device.device.ep_name = "DmlExecutionProvider" + session_options = MagicMock() + config = ORTGraphPipeConfig(ep_device=ep_device) + + with ( + patch("onnxruntime.SessionOptions", return_value=session_options), + patch("onnxruntime.InferenceSession") as inference_session, + patch("winml.modelkit.optim.pipes.graph.save_onnx"), + patch("winml.modelkit.session.lookup_device_spec", return_value=None), + pytest.raises(OptimizationError, match="was not activated"), + ): + inference_session.return_value.get_providers.return_value = [ + "CPUExecutionProvider" + ] + ORTGraphPipe().process(model, config) + def test_always_disabled_optimizers(self) -> None: """AttentionFusion and EmbedLayerNormFusion are ALWAYS disabled. diff --git a/tests/unit/optim/test_analysis.py b/tests/unit/optim/test_analysis.py index 6c778ee4f..2c9d1fd3b 100644 --- a/tests/unit/optim/test_analysis.py +++ b/tests/unit/optim/test_analysis.py @@ -12,6 +12,8 @@ from __future__ import annotations +from unittest.mock import MagicMock + import numpy as np from onnx import GraphProto, ModelProto, TensorProto, helper, numpy_helper @@ -30,6 +32,7 @@ _diff_initializers, _diff_nodes, ) +from winml.modelkit.optim.registry import CapabilityCategory # ============================================================================= @@ -418,6 +421,70 @@ def test_input_model_not_mutated(self) -> None: after_value = numpy_helper.to_array(model.graph.initializer[0]) np.testing.assert_array_equal(before_value, after_value) + def test_target_context_forwarded_and_ep_constraints_filtered( + self, monkeypatch, caplog + ) -> None: + dml_cap = BoolCapability( + name="dml-only", + ort_name=None, + description="DML-only probe", + category=CapabilityCategory.MISC, + ep_constraint=("DML",), + ) + cuda_cap = BoolCapability( + name="cuda-only", + ort_name=None, + description="CUDA-only probe", + category=CapabilityCategory.MISC, + ep_constraint=("CUDA",), + ) + cap_registry = {cap.name: cap for cap in (dml_cap, cuda_cap)} + captured_ep_devices = [] + started: list[str] = [] + completed: list[str] = [] + + class TargetAwarePipe: + name = "target-aware" + capabilities = cap_registry + + @classmethod + def build_config(cls, **kwargs): + captured_ep_devices.append(kwargs.get("ep_device")) + return kwargs + + @staticmethod + def process(model, config): + enabled = next( + (name for name in ("dml_only", "cuda_only") if config.get(name)), + None, + ) + if enabled is not None: + model.graph.node.append( + helper.make_node("Identity", ["z"], [f"{enabled}_output"]) + ) + return model + + ep_device = MagicMock() + ep_device.device.ep_name = "DmlExecutionProvider" + caplog.set_level("DEBUG", logger="winml.modelkit.optim.analysis") + monkeypatch.setattr("winml.modelkit.optim.pipes.PIPES", [TargetAwarePipe]) + monkeypatch.setattr("winml.modelkit.onnx.infer_shapes", lambda model: model) + + findings = analyze_model( + _benign_model(), + cap_registry, + ep_device=ep_device, + on_probe_start=started.append, + on_probe_complete=completed.append, + ) + + assert [finding.name for finding in findings] == ["dml-only"] + assert captured_ep_devices + assert all(value is ep_device for value in captured_ep_devices) + assert "Skipping capability 'cuda-only'" in caplog.text + assert started == ["dml-only", "cuda-only"] + assert completed == ["dml-only", "cuda-only"] + # ============================================================================= # OPTIMIZATION OUTPUT ITERATION