From c04228c25abe8b0113bb8962e70baf850ca1b750 Mon Sep 17 00:00:00 2001 From: hualxie Date: Thu, 30 Jul 2026 16:57:22 +0800 Subject: [PATCH 1/3] Add EP device targeting to optimize --- src/winml/modelkit/commands/optimize.py | 59 +++++++++++++++++++- src/winml/modelkit/optim/pipes/graph.py | 41 ++++++++++++-- tests/unit/commands/test_optimize_cli.py | 64 +++++++++++++++++++++- tests/unit/optim/pipes/test_pipe_config.py | 39 +++++++++++++ 4 files changed, 195 insertions(+), 8 deletions(-) diff --git a/src/winml/modelkit/commands/optimize.py b/src/winml/modelkit/commands/optimize.py index 0b5a7917b..297177200 100644 --- a/src/winml/modelkit/commands/optimize.py +++ b/src/winml/modelkit/commands/optimize.py @@ -290,6 +290,18 @@ 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, + 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", @@ -309,6 +321,8 @@ def optimize( model: Path | None, output: Path | None, overwrite: bool, + ep: str | None, + device: str | None, config: Path | None, verbose: int, quiet: bool, @@ -529,6 +543,36 @@ def optimize( if verbose: optimizer_kwargs["verbose"] = True + if ep is not None or device is not None: + from ..session import ( + DeviceNotFound, + EPDeviceTarget, + UnknownListingPick, + WinMLEPNotDiscovered, + WinMLEPRegistrationFailed, + WinMLEPRegistry, + resolve_device, + ) + + try: + target = resolve_device( + EPDeviceTarget(ep=ep or "auto", device=(device or "auto").lower()) + ) + ep_device = WinMLEPRegistry.instance().auto_device(target) + except ( + DeviceNotFound, + RuntimeError, + UnknownListingPick, + ValueError, + WinMLEPNotDiscovered, + WinMLEPRegistrationFailed, + ) as e: + raise click.UsageError(f"Could not resolve optimization target: {e}") from e + 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) @@ -543,10 +587,21 @@ 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 = "change" + node_info = ( + f"Nodes: {original_nodes} -> {optimized_nodes} " + f"({node_change:.1f}% {change_label})" + ) console.print(f"[dim]{node_info}[/dim]") except Exception as e: diff --git a/src/winml/modelkit/optim/pipes/graph.py b/src/winml/modelkit/optim/pipes/graph.py index 853cdb9ee..6ea5feeae 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,31 @@ 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(), + ) + 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: + _ = ort.InferenceSession( + str(input_file), sess_opts, providers=["CPUExecutionProvider"] + ) + else: + _ = ort.InferenceSession(str(input_file), sess_opts) except Exception as e: raise OptimizationError( f"ONNX Runtime optimization failed: {e}", diff --git a/tests/unit/commands/test_optimize_cli.py b/tests/unit/commands/test_optimize_cli.py index d5b6507aa..58586c1a2 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,56 @@ 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_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 + ) # ============================================================================= diff --git a/tests/unit/optim/pipes/test_pipe_config.py b/tests/unit/optim/pipes/test_pipe_config.py index 4e6093daf..3e46b6bc0 100644 --- a/tests/unit/optim/pipes/test_pipe_config.py +++ b/tests/unit/optim/pipes/test_pipe_config.py @@ -22,6 +22,7 @@ import logging from typing import TYPE_CHECKING +from unittest.mock import MagicMock, patch from winml.modelkit.optim.pipes import ( GRAPH_CAPABILITIES, @@ -131,6 +132,44 @@ 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) -> None: + model = MagicMock() + optimized_model = MagicMock() + ep_device = MagicMock() + handle = ep_device.device.ort_handle + 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.optim.pipes.graph.load_onnx", + return_value=optimized_model, + ), + patch("winml.modelkit.session.lookup_device_spec", return_value=None), + ): + 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 + def test_always_disabled_optimizers(self) -> None: """AttentionFusion and EmbedLayerNormFusion are ALWAYS disabled. From f060afed9e8dcc71ad86c728a81fcaeffc10aa9b Mon Sep 17 00:00:00 2001 From: hualxie Date: Thu, 30 Jul 2026 17:33:31 +0800 Subject: [PATCH 2/3] Make optimization checks target aware --- src/winml/modelkit/commands/optimize.py | 85 ++++++++++++++++-------- src/winml/modelkit/optim/analysis.py | 23 ++++++- src/winml/modelkit/utils/cli.py | 9 ++- tests/unit/commands/test_optimize_cli.py | 43 ++++++++++++ tests/unit/optim/test_analysis.py | 59 ++++++++++++++++ 5 files changed, 188 insertions(+), 31 deletions(-) diff --git a/src/winml/modelkit/commands/optimize.py b/src/winml/modelkit/commands/optimize.py index 297177200..8b1c01bd4 100644 --- a/src/winml/modelkit/commands/optimize.py +++ b/src/winml/modelkit/commands/optimize.py @@ -228,7 +228,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 @@ -238,6 +257,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 @@ -246,6 +267,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]" ) @@ -257,7 +280,7 @@ def _run_check_optim(model: Path, all_caps: dict[str, Any], verbose: bool) -> No "[dim](this can take a while on large models)[/dim]" ) with console.status("[bold]Analyzing...[/bold]", spinner="dots"): - findings = analyze_model(onnx_model, all_caps) + findings = analyze_model(onnx_model, all_caps, ep_device=ep_device) _render_check_optim(console, findings, verbose) @@ -294,6 +317,7 @@ def _run_check_optim(model: Path, all_caps: dict[str, Any], verbose: bool) -> No 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( @@ -477,9 +501,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 @@ -543,31 +595,8 @@ def optimize( if verbose: optimizer_kwargs["verbose"] = True - if ep is not None or device is not None: - from ..session import ( - DeviceNotFound, - EPDeviceTarget, - UnknownListingPick, - WinMLEPNotDiscovered, - WinMLEPRegistrationFailed, - WinMLEPRegistry, - resolve_device, - ) - - try: - target = resolve_device( - EPDeviceTarget(ep=ep or "auto", device=(device or "auto").lower()) - ) - ep_device = WinMLEPRegistry.instance().auto_device(target) - except ( - DeviceNotFound, - RuntimeError, - UnknownListingPick, - ValueError, - WinMLEPNotDiscovered, - WinMLEPRegistrationFailed, - ) as e: - raise click.UsageError(f"Could not resolve optimization target: {e}") from e + if ep_device is not None: + assert target is not None optimizer_kwargs["ep_device"] = ep_device console.print( f"[bold blue]Target:[/bold blue] {target.ep} on {target.device.upper()}" diff --git a/src/winml/modelkit/optim/analysis.py b/src/winml/modelkit/optim/analysis.py index 649a5d2d3..f0b3fcb11 100644 --- a/src/winml/modelkit/optim/analysis.py +++ b/src/winml/modelkit/optim/analysis.py @@ -283,6 +283,7 @@ def _run_pipe(pipe: Any, model: ModelProto, config: Any) -> ModelProto: def _iter_findings( model: ModelProto, capabilities: dict[str, CapabilityDef], + **optimizer_kwargs: Any, ) -> Iterator[tuple[CapabilityFinding, ModelProto]]: """Yield ``(finding, produced_model)`` for every applicable optimization. @@ -305,6 +306,7 @@ def _iter_findings( # 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: @@ -340,6 +342,14 @@ def _iter_findings( ] for cap_name, cap in probe_caps: + 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 = ep_device.device.ep_name + if not any(normalize_ep_name(name) == target_ep for name in cap.ep_constraint): + continue + # Enable only this capability (plus its dependencies) on top of the # all-defaults configuration. kebab = dict(kebab_defaults) @@ -350,6 +360,7 @@ def _iter_findings( 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) @@ -406,6 +417,7 @@ def _iter_findings( def analyze_model( model: ModelProto, capabilities: dict[str, CapabilityDef], + **optimizer_kwargs: Any, ) -> list[CapabilityFinding]: """Probe every applicable optimization capability against ``model``. @@ -418,17 +430,23 @@ 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``. Returns: Applicable findings in pipeline order, each naming the affected nodes and constants. """ - return [finding for finding, _ in _iter_findings(model, capabilities)] + return [ + finding + for finding, _ in _iter_findings(model, capabilities, **optimizer_kwargs) + ] 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. @@ -442,10 +460,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/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 58586c1a2..feef461ef 100644 --- a/tests/unit/commands/test_optimize_cli.py +++ b/tests/unit/commands/test_optimize_cli.py @@ -332,6 +332,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 + class TestConfigFile: """Config file loading and precedence.""" diff --git a/tests/unit/optim/test_analysis.py b/tests/unit/optim/test_analysis.py index a69856d17..1c7c59190 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 @@ -29,6 +31,7 @@ _diff_nodes, ) from winml.modelkit.optim.pipes import get_all_capabilities +from winml.modelkit.optim.registry import BoolCapability, CapabilityCategory # ============================================================================= @@ -403,6 +406,62 @@ 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 + ) -> 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 = [] + + 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" + 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, + ) + + 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) + # ============================================================================= # OPTIMIZATION OUTPUT ITERATION From ec7937e3b755af6ac9734471be23a70a284f27d1 Mon Sep 17 00:00:00 2001 From: hualxie Date: Fri, 31 Jul 2026 10:51:27 +0800 Subject: [PATCH 3/3] Address optimize target review feedback --- src/winml/modelkit/commands/optimize.py | 14 +++++----- src/winml/modelkit/optim/analysis.py | 8 +++++- src/winml/modelkit/optim/pipes/graph.py | 20 +++++++++++++-- tests/unit/commands/test_optimize_cli.py | 16 ++++++++++++ tests/unit/optim/pipes/test_pipe_config.py | 30 ++++++++++++++++++---- tests/unit/optim/test_analysis.py | 4 ++- 6 files changed, 77 insertions(+), 15 deletions(-) diff --git a/src/winml/modelkit/commands/optimize.py b/src/winml/modelkit/commands/optimize.py index 8b1c01bd4..5ef06c1ed 100644 --- a/src/winml/modelkit/commands/optimize.py +++ b/src/winml/modelkit/commands/optimize.py @@ -596,7 +596,8 @@ def optimize( optimizer_kwargs["verbose"] = True if ep_device is not None: - assert target 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()}" @@ -626,11 +627,12 @@ def optimize( elif optimized_nodes > original_nodes: change_label = "increase" else: - change_label = "change" - node_info = ( - f"Nodes: {original_nodes} -> {optimized_nodes} " - f"({node_change:.1f}% {change_label})" - ) + 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 f0b3fcb11..3d59adf98 100644 --- a/src/winml/modelkit/optim/analysis.py +++ b/src/winml/modelkit/optim/analysis.py @@ -346,8 +346,14 @@ def _iter_findings( if ep_device is not None and cap.ep_constraint is not None: from ..utils.constants import normalize_ep_name - target_ep = ep_device.device.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 diff --git a/src/winml/modelkit/optim/pipes/graph.py b/src/winml/modelkit/optim/pipes/graph.py index 6ea5feeae..4e5f4b3d1 100644 --- a/src/winml/modelkit/optim/pipes/graph.py +++ b/src/winml/modelkit/optim/pipes/graph.py @@ -597,11 +597,11 @@ def process(self, model: onnx.ModelProto, config: ORTGraphPipeConfig) -> onnx.Mo # Create session to trigger optimization try: if config.ep_device is None: - _ = ort.InferenceSession( + session = ort.InferenceSession( str(input_file), sess_opts, providers=["CPUExecutionProvider"] ) else: - _ = ort.InferenceSession(str(input_file), sess_opts) + session = ort.InferenceSession(str(input_file), sess_opts) except Exception as e: raise OptimizationError( f"ONNX Runtime optimization failed: {e}", @@ -613,6 +613,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/tests/unit/commands/test_optimize_cli.py b/tests/unit/commands/test_optimize_cli.py index feef461ef..6ad9d089b 100644 --- a/tests/unit/commands/test_optimize_cli.py +++ b/tests/unit/commands/test_optimize_cli.py @@ -220,6 +220,22 @@ def test_node_increase_reported(self, runner: CliRunner, tmp_path: Path) -> None 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: diff --git a/tests/unit/optim/pipes/test_pipe_config.py b/tests/unit/optim/pipes/test_pipe_config.py index 3e46b6bc0..424200190 100644 --- a/tests/unit/optim/pipes/test_pipe_config.py +++ b/tests/unit/optim/pipes/test_pipe_config.py @@ -21,21 +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) # ============================================================================= @@ -164,12 +162,34 @@ def test_process_binds_resolved_ep_device(self) -> None: ), 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 + 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 1c7c59190..667ff7f09 100644 --- a/tests/unit/optim/test_analysis.py +++ b/tests/unit/optim/test_analysis.py @@ -407,7 +407,7 @@ def test_input_model_not_mutated(self) -> None: np.testing.assert_array_equal(before_value, after_value) def test_target_context_forwarded_and_ep_constraints_filtered( - self, monkeypatch + self, monkeypatch, caplog ) -> None: dml_cap = BoolCapability( name="dml-only", @@ -449,6 +449,7 @@ def process(model, config): 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) @@ -461,6 +462,7 @@ def process(model, config): 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 # =============================================================================