From 258a3c76122c266b59566c48434370142f20d7d4 Mon Sep 17 00:00:00 2001 From: fangyangci Date: Thu, 30 Jul 2026 22:49:36 +0800 Subject: [PATCH] feat(analyze): use pattern rules for runtime support Extract and deduplicate pattern matches, resolve pattern rule support, and surface statuses through analyzer and CLI output. Improve unknown pattern, Gemm, RTR, and parquet query handling; remove obsolete subgraph runtime, debug logging, timing, and dead helper paths; add regression coverage. --- src/winml/modelkit/analyze/analyzer.py | 343 ++- .../modelkit/analyze/core/onnx_loader.py | 2 - .../analyze/core/pattern_extractor.py | 1929 ++++++++++++----- .../modelkit/analyze/core/runtime_checker.py | 282 +-- .../analyze/core/runtime_checker_query.py | 320 +-- .../modelkit/analyze/models/runtime_checks.py | 1 + .../runtime_checker/result_processor.py | 189 +- .../modelkit/analyze/utils/json_utils.py | 32 - src/winml/modelkit/analyze/utils/op_utils.py | 10 - .../modelkit/analyze/utils/rule_loader.py | 8 - src/winml/modelkit/cli.py | 2 +- src/winml/modelkit/commands/analyze.py | 537 +++-- src/winml/modelkit/pattern/base.py | 3 +- .../core/test_pattern_deduplication.py | 111 +- .../analyze/core/test_pattern_extractor.py | 569 +++-- .../unit/analyze/core/test_runtime_checker.py | 297 +-- .../test_runtime_checker_query_parquet.py | 41 + tests/unit/analyze/test_analyzer.py | 23 +- .../unit/analyze/test_static_analyzer_cli.py | 26 + 19 files changed, 2817 insertions(+), 1908 deletions(-) diff --git a/src/winml/modelkit/analyze/analyzer.py b/src/winml/modelkit/analyze/analyzer.py index e9768b47d..32a9c5aa9 100644 --- a/src/winml/modelkit/analyze/analyzer.py +++ b/src/winml/modelkit/analyze/analyzer.py @@ -20,20 +20,24 @@ from ..optim.config import WinMLOptimizationConfig from ..utils.constants import normalize_ep_name from .models.information import Information -from .models.output import RuntimeDebugSummaryEntry +from .models.output import ModelStats, RuntimeDebugSummaryEntry from .models.support_level import SupportLevel from .utils.timing_utils import make_timing_logger if TYPE_CHECKING: from collections.abc import Callable + from collections.abc import Mapping, Sequence import onnx + from onnx import NodeProto from ..utils.constants import EPName, EPNameOrAlias from .models.information import Action + from .models.onnx_model import ONNXModel from .models.output import AnalysisOutput from .models.runtime_checks import PatternRuntime, RuntimeTestResult + from ..pattern.match import PatternMatchResult @dataclass @@ -70,6 +74,13 @@ class LintResult: SupportLevel.SUPPORTED, ) +_PATTERN_STATUS_QUALITY: dict[str, int] = { + "unknow": 0, + "unsupported": 1, + "partial": 2, + "supported": 3, +} + def _normalize_case_indices_for_summary(case_indices: Any) -> list[Any] | None: """Normalize case_indices to JSON-friendly list values.""" @@ -89,6 +100,177 @@ def _iter_runtime_test_results(pattern_runtime: PatternRuntime) -> list[RuntimeT return results +def _candidate_to_supported_status(candidate: dict[str, Any] | None) -> str: + """Map compile/run candidate output to exported support status.""" + if not candidate: + return "unknow" + + if candidate.get("status") != "ok": + return "unknow" + + compile_ok = candidate.get("compile") + run_ok = candidate.get("run") + + if compile_ok is True and run_ok is True: + return "supported" + if compile_ok is False and run_ok is True: + return "partial" + return "unsupported" + + +def _pick_worst_status(statuses: list[str]) -> str: + """Pick worst status for one pattern group across all its instances.""" + if not statuses: + return "unknow" + return min(statuses, key=lambda status: _PATTERN_STATUS_QUALITY.get(status, 0)) + + +def _build_match_status_by_match_id( + merge_prep_entries: Sequence[Mapping[str, Any]], +) -> dict[str, str]: + """Build best available support status for each pattern match_id.""" + status_by_match_id: dict[str, str] = {} + for entry in merge_prep_entries: + match_id = str(entry.get("match_id", "")) + if not match_id: + continue + + raw_status = str(entry.get("support_status", "")).strip().lower() + if raw_status == "unknown": + raw_status = "unknow" + + if raw_status in _PATTERN_STATUS_QUALITY: + status = raw_status + else: + pattern_id = str(entry.get("pattern_id", "")) + candidates = entry.get("candidates", []) or [] + base_candidate = next( + ( + candidate + for candidate in candidates + if not bool(candidate.get("is_alternative", False)) + and str(candidate.get("pattern_id", "")) == pattern_id + ), + None, + ) + status = _candidate_to_supported_status(base_candidate) + + status_by_match_id[match_id] = status + + return status_by_match_id + + +def _build_pattern_status_by_node_key( + subgraph_patterns: list[PatternMatchResult], + merge_prep_entries: Sequence[Mapping[str, Any]], +) -> dict[str, str]: + """Build per-node pattern status map for matched-node runtime short-circuit.""" + status_by_match_id = _build_match_status_by_match_id(merge_prep_entries) + status_by_node_key: dict[str, str] = {} + + for pattern_match in subgraph_patterns: + status = status_by_match_id.get(pattern_match.match_id, "unknow") + for node_key in pattern_match.matched_node_keys: + status_by_node_key[node_key] = status + + return status_by_node_key + + +def _build_pattern_matching_summary( + subgraph_patterns: list[PatternMatchResult], + merge_prep_entries: Sequence[Mapping[str, Any]], +) -> dict[str, Any]: + """Build per-EP pattern summary payload for CLI rendering.""" + status_by_match_id = _build_match_status_by_match_id(merge_prep_entries) + + grouped: dict[str, dict[str, Any]] = {} + covered_node_keys: set[str] = set() + + for pattern_match in subgraph_patterns: + pattern_id = pattern_match.pattern.pattern_id + status = status_by_match_id.get(pattern_match.match_id, "unknow") + covered_node_keys.update(pattern_match.matched_node_keys) + + bucket = grouped.setdefault( + pattern_id, + { + "pattern_id": pattern_id, + "statuses": [], + "instances": 0, + "node_op_counts": {}, + }, + ) + bucket["statuses"].append(status) + bucket["instances"] += 1 + + skeleton_nodes = pattern_match.skeleton_match_result.matched_nodes + for node in skeleton_nodes: + op_type = node.op_type + node_op_counts: dict[str, int] = bucket["node_op_counts"] + node_op_counts[op_type] = node_op_counts.get(op_type, 0) + 1 + + patterns: list[dict[str, Any]] = [] + for pattern_id, bucket in grouped.items(): + instances = int(bucket["instances"]) + op_counts: dict[str, int] = bucket["node_op_counts"] + + node_breakdown: list[dict[str, Any]] = [] + for op_type, total_count in sorted(op_counts.items(), key=lambda item: (-item[1], item[0])): + per_instance_count = total_count // instances if instances > 0 and total_count % instances == 0 else None + node_breakdown.append( + { + "op_type": op_type, + "per_instance_count": per_instance_count, + "total_count": total_count, + } + ) + + total_child_nodes = sum(op_counts.values()) + patterns.append( + { + "pattern_id": pattern_id, + "status": _pick_worst_status(bucket["statuses"]), + "instances": instances, + "node_breakdown": node_breakdown, + "total_child_nodes": total_child_nodes, + } + ) + + patterns.sort(key=lambda item: (-int(item["instances"]), str(item["pattern_id"]))) + return { + "patterns": patterns, + "pattern_nodes_total": len(covered_node_keys), + } + + +def _build_operator_counts_excluding_pattern_nodes( + *, + operator_counts: Mapping[str, int], + onnx_model: "ONNXModel", + matched_node_keys: set[str], +) -> dict[str, int]: + """Subtract pattern-matched nodes from operator totals for OP CHECK display.""" + if not matched_node_keys: + return {str(op_type): int(count) for op_type, count in operator_counts.items() if int(count) > 0} + + matched_counts_by_op: dict[str, int] = {} + for node_key in matched_node_keys: + node = onnx_model.get_node_by_key(node_key) + if node is None: + continue + op_type = str(node.op_type) + matched_counts_by_op[op_type] = matched_counts_by_op.get(op_type, 0) + 1 + + adjusted_counts: dict[str, int] = {} + for op_type, total_count_raw in operator_counts.items(): + total_count = int(total_count_raw) + remaining = total_count - matched_counts_by_op.get(str(op_type), 0) + if remaining > 0: + adjusted_counts[str(op_type)] = remaining + + return adjusted_counts + + def _build_runtime_debug_details_summary( runtime_summary: dict[str, list[PatternRuntime]], ) -> dict[str, list[str] | dict[str, RuntimeDebugSummaryEntry]] | None: @@ -105,50 +287,47 @@ def _build_runtime_debug_details_summary( } unknown_nodes: set[str] = set() - for runtime_key in ("op_runtime_check_result", "subgraph_runtime_check_result"): - for pattern_runtime in runtime_summary.get(runtime_key, []): - for test_result in _iter_runtime_test_results(pattern_runtime): - level = test_result.classification + for pattern_runtime in runtime_summary.get("op_runtime_check_result", []): + for test_result in _iter_runtime_test_results(pattern_runtime): + level = test_result.classification - debug_details = test_result.debug_details - if not debug_details: - continue + debug_details = test_result.debug_details + if not debug_details: + continue - node_stable_key = debug_details.get("node_stable_key") - if not node_stable_key: - continue + node_stable_key = debug_details.get("node_stable_key") + if not node_stable_key: + continue - if level == SupportLevel.UNKNOWN: - # Unknown nodes carry no rule case data; record the - # de-duplicated node key only. - unknown_nodes.add(node_stable_key) - continue + if level == SupportLevel.UNKNOWN: + # Unknown nodes carry no rule case data; record the + # de-duplicated node key only. + unknown_nodes.add(node_stable_key) + continue - if level not in _RUNTIME_DEBUG_SUMMARY_LEVELS: - continue + if level not in _RUNTIME_DEBUG_SUMMARY_LEVELS: + continue - candidate_entry = RuntimeDebugSummaryEntry( - case_indices=_normalize_case_indices_for_summary( - debug_details.get("case_indices") - ), - table_path=debug_details.get("table_path"), - table_file=debug_details.get("table_file"), - ) + candidate_entry = RuntimeDebugSummaryEntry( + case_indices=_normalize_case_indices_for_summary(debug_details.get("case_indices")), + table_path=debug_details.get("table_path"), + table_file=debug_details.get("table_file"), + ) - level_bucket = leveled_summary[level.value] - existing_entry = level_bucket.get(node_stable_key) - if existing_entry is None: - level_bucket[node_stable_key] = candidate_entry - continue + level_bucket = leveled_summary[level.value] + existing_entry = level_bucket.get(node_stable_key) + if existing_entry is None: + level_bucket[node_stable_key] = candidate_entry + continue - if existing_entry.case_indices is None and candidate_entry.case_indices is not None: - existing_entry.case_indices = candidate_entry.case_indices + if existing_entry.case_indices is None and candidate_entry.case_indices is not None: + existing_entry.case_indices = candidate_entry.case_indices - if existing_entry.table_path is None and candidate_entry.table_path is not None: - existing_entry.table_path = candidate_entry.table_path + if existing_entry.table_path is None and candidate_entry.table_path is not None: + existing_entry.table_path = candidate_entry.table_path - if existing_entry.table_file is None and candidate_entry.table_file is not None: - existing_entry.table_file = candidate_entry.table_file + if existing_entry.table_file is None and candidate_entry.table_file is not None: + existing_entry.table_file = candidate_entry.table_file has_any_entry = bool(unknown_nodes) or any( leveled_summary[level.value] for level in _RUNTIME_DEBUG_SUMMARY_LEVELS @@ -174,13 +353,16 @@ class AnalysisResult: def __init__( self, output: AnalysisOutput, + pattern_matching_by_ep: dict[str, dict[str, Any]] | None = None, ) -> None: """Initialize analysis result. Args: output: The analysis output + pattern_matching_by_ep: Per-EP pattern summary payload for CLI rendering. """ self.output = output + self.pattern_matching_by_ep: dict[str, dict[str, Any]] = pattern_matching_by_ep or {} def __repr__(self) -> str: """String representation of analysis result.""" @@ -604,6 +786,9 @@ def analyze( save_node_types: set[str] | None = None, on_node_result: Callable | None = None, on_ep_start: Callable | None = None, + on_pattern_query_start: Callable | None = None, + on_pattern_query_result: Callable | None = None, + on_pattern_summary_ready: Callable | None = None, ) -> AnalysisResult: """Analyze ONNX model for runtime support. @@ -713,6 +898,9 @@ def analyze( save_node_types=save_node_types, on_node_result=on_node_result, on_ep_start=on_ep_start, + on_pattern_query_start=on_pattern_query_start, + on_pattern_query_result=on_pattern_query_result, + on_pattern_summary_ready=on_pattern_summary_ready, ) delegate_ms = int((time.perf_counter() - delegate_start) * 1000) _log_timing( @@ -739,6 +927,9 @@ def analyze_from_proto( save_node_types: set[str] | None = None, on_node_result: Callable | None = None, on_ep_start: Callable | None = None, + on_pattern_query_start: Callable | None = None, + on_pattern_query_result: Callable | None = None, + on_pattern_summary_ready: Callable | None = None, ) -> AnalysisResult: """Analyze ONNX model from ModelProto object. @@ -837,13 +1028,14 @@ def analyze_from_proto( object.__setattr__(onnx_model, "model_path", model_path) pattern_extractor = PatternExtractor(onnx_model, htp_metadata_path=htp_metadata_path) - extraction_result = pattern_extractor.summary() - - metadata = extraction_result["summary"] - pattern_matches = extraction_result["subgraph_patterns"] - logger.info("Extracted %d patterns", len(pattern_matches)) + metadata = pattern_extractor.model_summary() extraction_ms = int((time.perf_counter() - extraction_start) * 1000) + # Keep subgraph runtime aggregation disabled for now. Pattern extraction + # still drives per-EP node skip sets and pattern UI payloads. + pattern_matching_by_ep: dict[str, dict[str, Any]] = {} + pattern_count_for_timing = 0 + # Step 2: Check runtime support for each EP check_op_results: dict[EPName, list[PatternRuntime]] = {} information_list: dict[EPName, list[Information]] = {} @@ -854,17 +1046,74 @@ def analyze_from_proto( ep_info_timing: dict[str, int] = {} for current_ep in eps_to_analyze: logger.info("Checking runtime support for %s...", current_ep) + + def _on_pattern_query_start_for_ep(pattern_counts: Mapping[str, int]) -> None: + if on_pattern_query_start is None: + return + try: + on_pattern_query_start(current_ep, dict(pattern_counts)) + except Exception: + logger.debug("on_pattern_query_start callback failed", exc_info=True) + + def _on_pattern_query_result_for_ep(pattern_id: str, support_status: str) -> None: + if on_pattern_query_result is None: + return + try: + on_pattern_query_result(current_ep, pattern_id, support_status) + except Exception: + logger.debug("on_pattern_query_result callback failed", exc_info=True) + + ep_pattern_summary = pattern_extractor.summary( + ep=current_ep, + device=device_to_use, + for_debug=for_debug, + on_pattern_query_start=_on_pattern_query_start_for_ep, + on_pattern_query_result=_on_pattern_query_result_for_ep, + ) + # In single-EP mode, keep metadata pattern counts aligned with the + # selected EP summary while still avoiding duplicate upfront extraction. + # Also tolerate minimal test doubles that don't stub model_summary(). + if ep_normalized is not None or not isinstance(metadata, ModelStats): + metadata = ep_pattern_summary["summary"] + + ep_subgraph_patterns = ep_pattern_summary["subgraph_patterns"] + ep_merge_prep = ep_pattern_summary.get("merge_prep", []) + if not pattern_matching_by_ep: + pattern_count_for_timing = len(ep_subgraph_patterns) + + pattern_status_by_node_key = _build_pattern_status_by_node_key( + ep_subgraph_patterns, + ep_merge_prep, + ) + ep_pattern_payload = _build_pattern_matching_summary( + ep_subgraph_patterns, + ep_merge_prep, + ) + pattern_matching_by_ep[current_ep] = ep_pattern_payload + + if on_pattern_summary_ready is not None: + try: + on_pattern_summary_ready(current_ep, ep_pattern_payload) + except Exception: + logger.debug("on_pattern_summary_ready callback failed", exc_info=True) + if on_ep_start: try: - on_ep_start(current_ep, metadata.operator_counts) + op_counts_for_display = _build_operator_counts_excluding_pattern_nodes( + operator_counts=ep_pattern_summary["summary"].operator_counts, + onnx_model=onnx_model, + matched_node_keys=set(pattern_status_by_node_key), + ) + on_ep_start(current_ep, op_counts_for_display) except Exception: logger.debug("on_ep_start callback failed", exc_info=True) + runtime_summary_start = time.perf_counter() runtime_checker = RuntimeChecker( ep=current_ep, device=device_to_use, model=onnx_model, - patterns=pattern_matches, + pattern_matched_node_status_by_key=pattern_status_by_node_key, ) # TODO: add VitisAIExecutionProvider back once non-QDQ # data is ready, and run_unknown_op is supported for QDQ ops @@ -873,7 +1122,6 @@ def analyze_from_proto( run_unknown_op_for_ep = False runtime_summary = runtime_checker.summary( - patterns=pattern_matches, for_debug=for_debug, run_unknown_op=run_unknown_op_for_ep, save_node_types=save_node_types, @@ -889,7 +1137,6 @@ def analyze_from_proto( # Convert runtime summary to expected format op_results_list = runtime_summary.get("op_runtime_check_result", []) - subgraph_results_list = runtime_summary.get("subgraph_runtime_check_result", []) check_op_results[current_ep] = op_results_list # Use EP name as key @@ -901,7 +1148,7 @@ def analyze_from_proto( information_start = time.perf_counter() engine = self.information_engine_cls( op_runtime_results=op_results_list, - subgraph_runtime_results=subgraph_results_list, + subgraph_runtime_results=[], ep=current_ep, model=onnx_model, device=device_to_use, @@ -932,7 +1179,7 @@ def analyze_from_proto( ep=ep_normalized, device=device_to_use, eps=len(eps_to_analyze), - patterns=len(pattern_matches), + patterns=pattern_count_for_timing, extraction_ms=extraction_ms, aggregate_ms=aggregate_ms, runtime_ms_by_ep=ep_runtime_timing, @@ -941,7 +1188,7 @@ def analyze_from_proto( ) logger.info("Analysis complete") - return AnalysisResult(output=output) + return AnalysisResult(output=output, pattern_matching_by_ep=pattern_matching_by_ep) # ============================================================================= diff --git a/src/winml/modelkit/analyze/core/onnx_loader.py b/src/winml/modelkit/analyze/core/onnx_loader.py index 7794f00d0..8f111ca69 100644 --- a/src/winml/modelkit/analyze/core/onnx_loader.py +++ b/src/winml/modelkit/analyze/core/onnx_loader.py @@ -198,8 +198,6 @@ def validate(model_proto: onnx.ModelProto) -> None: if not model_proto.graph.node: raise ValueError("Model graph has no nodes") - # Skip strict ONNX validation to allow custom attributes like hierarchy_tag - # The model structure is still validated by checking for non-empty graph logger.debug("Skipping strict ONNX validation to allow custom attributes") def extract_metadata(self, detected_pattern_count: dict[str, int] | None = None) -> ModelStats: diff --git a/src/winml/modelkit/analyze/core/pattern_extractor.py b/src/winml/modelkit/analyze/core/pattern_extractor.py index 43d188c8d..491dee8d3 100644 --- a/src/winml/modelkit/analyze/core/pattern_extractor.py +++ b/src/winml/modelkit/analyze/core/pattern_extractor.py @@ -9,35 +9,97 @@ from __future__ import annotations +import copy +import hashlib +import json import logging +import re import time -from typing import TYPE_CHECKING, TypedDict, cast +from pathlib import Path +from typing import TYPE_CHECKING, Any, TypedDict, cast -from ...pattern.base import InvalidPatternMatcherModelError, PatternMatcher -from ...pattern.config import UnifiedPatternConfig +import numpy as np + +from ...onnx import ONNXDomain +from ...pattern.base import InvalidPatternMatcherModelError, PatternMatcher, PatternMismatchedError +from ...pattern.config import PatternConfig, UnifiedPatternConfig from ..models.onnx_model import ModelTag, ONNXModel from ..models.output import extract_model_stats +from ..utils.model_utils import encode_rule_condition_value_for_parquet, make_hashable +from ..utils.rule_loader import get_runtime_rules_debug_search_dirs, get_runtime_rules_search_dirs from ..utils.timing_utils import make_timing_logger if TYPE_CHECKING: + from collections.abc import Callable, Mapping + import onnx - from winml.modelkit.pattern.match import PatternMatchResult - from winml.modelkit.pattern.models import SubgraphPattern + from ...pattern.base import Pattern + from ...pattern.match import PatternMatchResult + from ...utils.constants import EPNameOrAlias + from ..models.ihv_type import IHVType from ..models.output import ModelStats +class PatternSourceStat(TypedDict): + """Per-source skeleton extraction stats for debug reporting.""" + + source: str + cache_hit: bool + pattern_class_count: int + match_count: int + elapsed_ms: int + + class PatternSummary(TypedDict): """Type definition for pattern analysis summary.""" summary: ModelStats subgraph_patterns: list[PatternMatchResult] - - -# Type alias for HTP metadata structure -HTPMetadata = dict[str, dict[str, str] | dict[str, object]] + subgraph_patterns_by_source: dict[str, dict[str, list[PatternMatchResult]]] + source_stats: list[PatternSourceStat] + merge_prep: list["PatternMergePrepEntry"] + model_signature: str + + +class PatternRuleCompileRunResult(TypedDict): + """Rule-table compile/run snapshot for one pattern candidate.""" + + pattern_class: str + pattern_id: str + is_alternative: bool + status: str + mismatch_error: str | None + compile: bool | None + run: bool | None + row_count: int + table_file: str | None + table_path: str | None + domain: str | None + opset_version: int | None + compile_true_rows: int + run_true_rows: int + case_indices: list[Any] | None + query_condition_count: int + query_condition_keys: list[str] + debug_details: dict[str, Any] | None + + +class PatternMergePrepEntry(TypedDict): + """Derived metadata used by upcoming pattern merge/dedup stage.""" + + source: str + pattern_class: str + pattern_id: str + match_count: int + match_index: int + match_id: str + matched_node_keys: list[str] + support_status: str + alternatives: list[dict[str, Any]] + candidates: list[PatternRuleCompileRunResult] logger = logging.getLogger(__name__) _log_timing = make_timing_logger(logger) @@ -58,6 +120,19 @@ class PatternExtractor: model: ONNX model to analyze (ONNXModel) """ + # In-memory per-process caches. + # - rules cache: source key -> loaded skeleton Pattern instances + # - match cache: (model signature, source key) -> grouped PatternMatchResult + # - merge prep cache: (model signature, ep, device, debug flag) -> merge prep entries + _RULES_PATTERN_CACHE: dict[str, list[Pattern]] = {} + _MATCH_CACHE: dict[tuple[str, str], dict[str, list[PatternMatchResult]]] = {} + _DEDUPED_MATCH_CACHE: dict[ + tuple[str, str], + tuple[dict[str, dict[str, list[PatternMatchResult]]], list[PatternMatchResult]], + ] = {} + _MERGE_PREP_CACHE: dict[tuple[str, str, str, bool], list[PatternMergePrepEntry]] = {} + _VALID_EP_DEVICE_PAIRS_CACHE: set[tuple[str, str]] | None = None + def __init__(self, model: ONNXModel, htp_metadata_path: str | None = None) -> None: """Initialize pattern extractor. @@ -72,8 +147,10 @@ def __init__(self, model: ONNXModel, htp_metadata_path: str | None = None) -> No raise TypeError(f"Expected ONNXModel, got {type(model)}") self._model = model - self._htp_metadata_path = htp_metadata_path - self._htp_metadata: HTPMetadata | None = None + self._query_condition_build_cache: dict[ + tuple[str, str, tuple[tuple[str, int], ...]], + tuple[dict[str, Any], Any], + ] = {} logger.info( "Initialized PatternExtractor for model: %s", @@ -88,585 +165,1449 @@ def model(self) -> ONNXModel: """The ONNX model being analyzed.""" return self._model - def _load_htp_metadata(self) -> HTPMetadata: - """Load HTP metadata from JSON file. + def _compute_model_signature(self) -> str: + """Build a stable in-process signature for cache keys.""" + model_path = self._model.model_path + if model_path and model_path != "": + path = Path(model_path) + if path.exists(): + stat = path.stat() + return f"{path.resolve()}|{stat.st_size}|{stat.st_mtime_ns}" + + # Fallback for in-memory models or missing paths. + model_bytes = self._model.get_model().SerializeToString() + digest = hashlib.sha1(model_bytes).hexdigest() + return f"in_memory:{digest}" + + @staticmethod + def _ihv_to_rules_key(ihv_type: IHVType) -> str | None: + """Map IHV enum to rules filename stem.""" + mapping = { + "QC": "qnn", + "INTEL": "openvino", + "AMD": "quark", + "NVIDIA": "nvidia", + "MICROSOFT": "microsoft", + } + return mapping.get(ihv_type.name) - Returns: - Dictionary containing HTP metadata + def _resolve_sources_for_ep(self, ep: EPNameOrAlias | None) -> list[str]: + """Return extraction sources for the target EP. - Raises: - FileNotFoundError: If metadata file doesn't exist - ValueError: If JSON is invalid + The new flow keeps default and IHV-specific extraction independent. """ - if self._htp_metadata is not None: - return self._htp_metadata + sources = ["default"] + if ep is None: + return sources + + from ..utils import infer_ihv_from_ep_name + from ..models.ihv_type import IHVType + + ihv_type = infer_ihv_from_ep_name(ep) + if ihv_type is IHVType.UNKNOWN: + return sources + + rules_key = self._ihv_to_rules_key(ihv_type) + if rules_key and self._rules_file_for_source(rules_key).exists(): + sources.append(rules_key) + return sources + + @staticmethod + def _rules_dir() -> Path: + """Return the pattern rules directory.""" + # .../modelkit/analyze/core/pattern_extractor.py -> .../modelkit/pattern/rules + return Path(__file__).resolve().parents[2] / "pattern" / "rules" + + def _rules_file_for_source(self, source: str) -> Path: + """Return rules JSON path for a source key.""" + return self._rules_dir() / f"{source}.json" + + @staticmethod + def _available_providers_config_path() -> Path: + """Return bundled EP/device validity mapping JSON path.""" + return ( + Path(__file__).resolve().parents[1] + / "utils" + / "avalizble_ep_device_ops" + / "avaliable_providers.json" + ) - if not self._htp_metadata_path: - logger.debug("No HTP metadata path provided") - return {} + @classmethod + def _load_valid_ep_device_pairs(cls) -> set[tuple[str, str]]: + """Load and cache valid EP/device pairs from provider config.""" + if cls._VALID_EP_DEVICE_PAIRS_CACHE is not None: + return cls._VALID_EP_DEVICE_PAIRS_CACHE - import json - from pathlib import Path + valid_pairs: set[tuple[str, str]] = set() + config_path = cls._available_providers_config_path() + try: + payload = json.loads(config_path.read_text(encoding="utf-8")) + except (OSError, json.JSONDecodeError): + logger.warning( + "Failed to load available providers config: %s", + config_path, + exc_info=True, + ) + cls._VALID_EP_DEVICE_PAIRS_CACHE = valid_pairs + return valid_pairs - metadata_path = Path(self._htp_metadata_path) - if not metadata_path.exists(): - raise FileNotFoundError(f"HTP metadata file not found: {self._htp_metadata_path}") + if not isinstance(payload, dict): + cls._VALID_EP_DEVICE_PAIRS_CACHE = valid_pairs + return valid_pairs - logger.info("Loading HTP metadata from: %s", self._htp_metadata_path) + for ep_name, ep_payload in payload.items(): + if not isinstance(ep_name, str) or not isinstance(ep_payload, dict): + continue - try: - with metadata_path.open(encoding="utf-8") as f: - self._htp_metadata = json.load(f) - logger.info("Successfully loaded HTP metadata") - return self._htp_metadata - except json.JSONDecodeError as e: - raise ValueError(f"Invalid JSON in HTP metadata file: {e}") from e - - def summary(self) -> PatternSummary: - """Generate comprehensive pattern analysis summary. + devices_payload = ep_payload.get("devices") + if not isinstance(devices_payload, dict): + continue - Returns: - PatternSummary with keys: - - summary: ModelStats (from model_summary()) - - subgraph_patterns: List[PatternMatchResult] (from extract_subgraph_patterns()) - """ - logger.info("Generating pattern analysis summary") - total_start = time.perf_counter() + for device_name, device_payload in devices_payload.items(): + if not isinstance(device_name, str) or not isinstance(device_payload, dict): + continue + if bool(device_payload.get("valid", False)): + valid_pairs.add((ep_name, device_name.upper())) - # Extract subgraph patterns - subgraph_start = time.perf_counter() - subgraph_patterns = self.extract_subgraph_patterns() - subgraph_ms = int((time.perf_counter() - subgraph_start) * 1000) + cls._VALID_EP_DEVICE_PAIRS_CACHE = valid_pairs + return valid_pairs - # Build pattern count dict: pattern_id -> count - count_dict_start = time.perf_counter() - pattern_count_dict: dict[str, int] = {} - for pattern_match in subgraph_patterns: - pattern_id = pattern_match.pattern.pattern_id - pattern_count_dict[pattern_id] = pattern_count_dict.get(pattern_id, 0) + 1 - count_dict_ms = int((time.perf_counter() - count_dict_start) * 1000) + def _is_valid_parquet_lookup_target(self, ep_name: str, device: str) -> bool: + """Return True when parquet lookup should run for this EP/device pair.""" + valid_pairs = self._load_valid_ep_device_pairs() + if not valid_pairs: + return False + return (ep_name, device.upper()) in valid_pairs + + def _load_skeleton_patterns_for_source(self, source: str) -> list[Pattern]: + """Load skeleton pattern instances for one source, with in-memory cache.""" + cached = self._RULES_PATTERN_CACHE.get(source) + if cached is not None: + return cached + + patterns: list[Pattern] = [] + if source == "default": + cfg = UnifiedPatternConfig(ihv_type="default") + patterns = cfg.get_skeleton_patterns() + self._RULES_PATTERN_CACHE[source] = patterns + return patterns + + rules_file = self._rules_file_for_source(source) + if not rules_file.exists(): + self._RULES_PATTERN_CACHE[source] = [] + return [] - # Generate model summary with pattern count dict - model_summary_start = time.perf_counter() - metadata = self.model_summary(detected_pattern_count=pattern_count_dict) - model_summary_ms = int((time.perf_counter() - model_summary_start) * 1000) + try: + with rules_file.open(encoding="utf-8") as f: + source_cfg = json.load(f) + except (OSError, json.JSONDecodeError): + logger.warning("Failed to load source rules config: %s", rules_file, exc_info=True) + self._RULES_PATTERN_CACHE[source] = [] + return [] - _log_timing( - "pattern_extractor.summary", - model=self._model.model_path, - detected_subgraph_patterns=len(subgraph_patterns), - unique_pattern_ids=len(pattern_count_dict), - extract_subgraph_ms=subgraph_ms, - build_count_dict_ms=count_dict_ms, - model_summary_ms=model_summary_ms, - total_ms=int((time.perf_counter() - total_start) * 1000), - ) + for entry in source_cfg.get("SkeletonPatternRules", []): + if not entry.get("enabled", False): + continue + try: + pattern_cfg = PatternConfig( + pattern_id=entry["pattern_id"], + pattern_class=entry["pattern_class"], + module=entry["module"], + enabled=bool(entry["enabled"]), + description=entry.get("description"), + alternatives=[], + ) + patterns.append(pattern_cfg.load_pattern()) + except Exception: + logger.warning( + "Failed to load skeleton pattern from %s for source '%s': %s", + rules_file, + source, + entry.get("pattern_class", ""), + exc_info=True, + ) + + self._RULES_PATTERN_CACHE[source] = patterns + return patterns + + def _extract_skeleton_matches_for_source( + self, + *, + source: str, + model_signature: str, + ) -> tuple[dict[str, list[PatternMatchResult]], PatternSourceStat]: + """Extract skeleton matches for one source with model+source cache key.""" + cache_key = (model_signature, source) + start = time.perf_counter() + + cached = self._MATCH_CACHE.get(cache_key) + if cached is not None: + elapsed_ms = int((time.perf_counter() - start) * 1000) + hit_stat: PatternSourceStat = { + "source": source, + "cache_hit": True, + "pattern_class_count": len(cached), + "match_count": sum(len(v) for v in cached.values()), + "elapsed_ms": elapsed_ms, + } + return {k: list(v) for k, v in cached.items()}, hit_stat + + grouped: dict[str, list[PatternMatchResult]] = {} + pattern_instances = self._load_skeleton_patterns_for_source(source) + if pattern_instances: + model_proto = self._model.get_model() + try: + matcher = PatternMatcher(model_proto, model_path=self._model.model_path) + except InvalidPatternMatcherModelError as e: + logger.warning("Model validation failed for pattern matching: %s", str(e)) + self._model.model_tags[ModelTag(e.error_tag)] = str(e) + matcher = None + + if matcher is not None: + for pattern in pattern_instances: + matcher.register_pattern(pattern) + + matches = matcher.match() + for match in matches: + # Keep explicit source for debug attribution. + match.attributes["source"] = source + pattern_class = match.pattern.__class__.__name__ + grouped.setdefault(pattern_class, []).append(match) + + self._MATCH_CACHE[cache_key] = grouped + elapsed_ms = int((time.perf_counter() - start) * 1000) + miss_stat: PatternSourceStat = { + "source": source, + "cache_hit": False, + "pattern_class_count": len(grouped), + "match_count": sum(len(v) for v in grouped.values()), + "elapsed_ms": elapsed_ms, + } + return {k: list(v) for k, v in grouped.items()}, miss_stat + @staticmethod + def _copy_grouped_matches( + grouped: dict[str, dict[str, list[PatternMatchResult]]], + ) -> dict[str, dict[str, list[PatternMatchResult]]]: + """Shallow-copy grouped-match containers while reusing match objects.""" return { - "summary": metadata, - "subgraph_patterns": subgraph_patterns, + source: { + pattern_class: list(matches) + for pattern_class, matches in source_group.items() + } + for source, source_group in grouped.items() } - def extract_subgraph_patterns(self) -> list[PatternMatchResult]: - """Extract subgraph patterns from model. + @staticmethod + def _cache_key_for_ep_dedup(ep: EPNameOrAlias | None) -> str: + """Build cache key component for EP-scoped dedup results.""" + if ep is None: + return "__default__" + return str(ep) - Subgraph patterns represent multi-operator fusion opportunities - (e.g., GELU, LayerNorm, Attention). + def _ordered_sources_for_ep_dedup( + self, + *, + sources: list[str], + ep: EPNameOrAlias | None, + ) -> list[str]: + """Return source traversal order with EP-specific source first when available.""" + if ep is None: + return list(sources) - Returns: - List of PatternMatchResult objects - - Process: - 1. Load subgraph pattern definitions via get_subgraph_patterns() - 2. For each pattern, match against model graph - 3. For each match, create PatternMatchResult with node_topology mapping - 4. Return all detected subgraph patterns - - Note: - - Pattern ID format: SUBGRAPH/ - - node_topology uses pattern-defined slot names as keys - - Actual node names from the model graph as values + from ..utils import infer_ihv_from_ep_name + from ..models.ihv_type import IHVType + + ihv_type = infer_ihv_from_ep_name(ep) + if ihv_type is IHVType.UNKNOWN: + return list(sources) + + ep_source = self._ihv_to_rules_key(ihv_type) + if not ep_source or ep_source not in sources: + return list(sources) + + return [ep_source] + [source for source in sources if source != ep_source] + + def _dedup_grouped_matches_for_ep( + self, + *, + subgraph_patterns_by_source: dict[str, dict[str, list[PatternMatchResult]]], + sources: list[str], + model_signature: str, + ep: EPNameOrAlias | None, + ) -> tuple[dict[str, dict[str, list[PatternMatchResult]]], list[PatternMatchResult]]: + """Deduplicate matches by node key with EP-source traversal priority. + + Traversal order is EP cache first (when present), then default cache. + Any pattern match touching a previously seen node key is filtered out. + Results are cached by (model signature, EP) so same EP with different + devices reuses dedup output. """ - logger.info("Extracting subgraph patterns from model") - total_start = time.perf_counter() + cache_key = (model_signature, self._cache_key_for_ep_dedup(ep)) + cached = self._DEDUPED_MATCH_CACHE.get(cache_key) + if cached is not None: + cached_grouped, cached_flat = cached + return self._copy_grouped_matches(cached_grouped), list(cached_flat) - # Get available subgraph pattern definitions - get_pattern_defs_start = time.perf_counter() - pattern_defs = self.get_subgraph_patterns() - get_pattern_defs_ms = int((time.perf_counter() - get_pattern_defs_start) * 1000) - - # Match patterns against model graph - detected_matches: list[PatternMatchResult] = [] - metadata_tag_match_start = time.perf_counter() - - for pattern_def in pattern_defs: - # Try HTP metadata-based matching first if available - if self._htp_metadata_path: - htp_matches = self._match_subgraph_pattern_from_htp_metadata(pattern_def) - if htp_matches: - detected_matches.extend(htp_matches) - continue + ordered_sources = self._ordered_sources_for_ep_dedup(sources=sources, ep=ep) - # Fall back to hierarchy_tag attribute-based matching - matches = self._match_subgraph_pattern_from_model_tags(pattern_def) - detected_matches.extend(matches) - metadata_tag_match_ms = int((time.perf_counter() - metadata_tag_match_start) * 1000) - - # Use PatternMatcher for skeleton-based pattern detection - logger.info("Using PatternMatcher for skeleton-based pattern detection") - pattern_matcher_start = time.perf_counter() - pattern_matcher_matches = self.extract_subgraph_patterns_with_pattern_matcher() - pattern_matcher_ms = int((time.perf_counter() - pattern_matcher_start) * 1000) - - # Deduplicate PatternMatcher results against existing matches - # Priority: HTP metadata > hierarchy_tag > PatternMatcher - # Collect node sets from existing matches (from HTP/tag) - dedup_start = time.perf_counter() - existing_node_sets: set[frozenset[str]] = { - frozenset(match.matched_nodes) for match in detected_matches - } + seen_node_keys: set[str] = set() + deduped_grouped: dict[str, dict[str, list[PatternMatchResult]]] = {} + deduped_flat: list[PatternMatchResult] = [] - # Filter PatternMatcher matches to exclude duplicates - filtered_matcher_matches: list[PatternMatchResult] = [] - for match in pattern_matcher_matches: - node_names = frozenset(match.matched_nodes) - if node_names not in existing_node_sets: - filtered_matcher_matches.append(match) - else: - # Log first few nodes (sorted for consistency) - sample_nodes = sorted(node_names)[:3] - logger.debug( - "Skipping PatternMatcher match with duplicate nodes: %s (pattern: %s)", - sample_nodes, - match.pattern_id, - ) + for source in ordered_sources: + source_group = subgraph_patterns_by_source.get(source, {}) + kept_by_pattern_class: dict[str, list[PatternMatchResult]] = {} - dropped_count = len(pattern_matcher_matches) - len(filtered_matcher_matches) - if dropped_count > 0: - logger.info( - "Dropped %d PatternMatcher matches that duplicate existing matches (from HTP/tag)", - dropped_count, - ) + for pattern_class, matches in source_group.items(): + kept_matches: list[PatternMatchResult] = [] - # Add filtered PatternMatcher matches - detected_matches.extend(filtered_matcher_matches) - dedup_ms = int((time.perf_counter() - dedup_start) * 1000) + for pattern_match in matches: + node_keys = list(pattern_match.matched_node_keys) + if any(node_key in seen_node_keys for node_key in node_keys): + continue - logger.info( - "Detected %d total subgraph pattern matches (including %d unique from PatternMatcher)", - len(detected_matches), - len(filtered_matcher_matches), - ) - _log_timing( - "pattern_extractor.extract_subgraph_patterns", - model=self._model.model_path, - pattern_defs=len(pattern_defs), - matches_before_matcher=len(existing_node_sets), - matcher_matches=len(pattern_matcher_matches), - matcher_unique_added=len(filtered_matcher_matches), - matcher_dropped_as_duplicate=dropped_count, - get_pattern_defs_ms=get_pattern_defs_ms, - metadata_tag_match_ms=metadata_tag_match_ms, - pattern_matcher_ms=pattern_matcher_ms, - dedup_ms=dedup_ms, - total_ms=int((time.perf_counter() - total_start) * 1000), - ) - return detected_matches + seen_node_keys.update(node_keys) + kept_matches.append(pattern_match) + deduped_flat.append(pattern_match) - def extract_subgraph_patterns_with_pattern_matcher(self) -> list[PatternMatchResult]: - """Extract subgraph patterns using PatternMatcher. + if kept_matches: + kept_by_pattern_class[pattern_class] = kept_matches - This method uses the PatternMatcher class to perform skeleton-based - pattern matching against registered patterns. + if kept_by_pattern_class: + deduped_grouped[source] = kept_by_pattern_class - Returns: - List of PatternMatchResult objects + self._DEDUPED_MATCH_CACHE[cache_key] = ( + self._copy_grouped_matches(deduped_grouped), + list(deduped_flat), + ) + return deduped_grouped, deduped_flat - Process: - 1. Create PatternMatcher instance with the model - 2. Load and register pattern instances from UnifiedPatternConfig - 3. Call matcher.match() to get PatternMatchResult objects - 4. Return all detected pattern matches - """ - logger.info("Extracting subgraph patterns using PatternMatcher") - total_start = time.perf_counter() + def _domain_and_target_opset_for_pattern( + self, + pattern: Pattern, + model_opsets: dict[ONNXDomain, int], + ) -> tuple[str, int]: + """Infer preferred domain/opset for locating pattern-level rule parquet files.""" + skeleton = pattern.get_skeleton() + if not skeleton.node_domains: + default_opset = model_opsets.get(ONNXDomain.AI_ONNX, 1) + return ONNXDomain.AI_ONNX.value, default_opset + + preferred_domain = skeleton.node_domains[0] + target_opset = model_opsets.get( + preferred_domain, + model_opsets.get(ONNXDomain.AI_ONNX, 1), + ) + return preferred_domain.value, target_opset + + @staticmethod + def _parse_pattern_rule_filename( + filename: str, + *, + pattern_class: str, + ep_name: str, + device: str, + ) -> tuple[str, int] | None: + """Parse `____opset.parquet` style names.""" + prefix = f"{pattern_class}_{ep_name}_{device.upper()}_" + if not filename.startswith(prefix): + return None - # Get model proto for PatternMatcher - get_model_start = time.perf_counter() - model_proto = self._model.get_model() - get_model_ms = int((time.perf_counter() - get_model_start) * 1000) + suffix = filename[len(prefix) :] + match = re.match(r"(?P.+)_opset(?P\d+)(?:_qdq)?\.parquet$", suffix) + if match is None: + return None - # Create PatternMatcher instance - may raise InvalidPatternMatcherModelError - try: - matcher_init_start = time.perf_counter() - matcher = PatternMatcher(model_proto, model_path=self._model.model_path) - matcher_init_ms = int((time.perf_counter() - matcher_init_start) * 1000) - except InvalidPatternMatcherModelError as e: - # Model is invalid for pattern matching (e.g., nodes with empty names) - logger.warning("Model validation failed for pattern matching: %s", str(e)) - # Mark model with the exception's associated tag and error message - self._model.model_tags[ModelTag(e.error_tag)] = str(e) - _log_timing( - "pattern_extractor.pattern_matcher", - model=self._model.model_path, - failed=True, - error_tag=e.error_tag, - get_model_ms=get_model_ms, - total_ms=int((time.perf_counter() - total_start) * 1000), - ) - return [] + return match.group("domain"), int(match.group("opset")) - # Register patterns from the unified pattern config - load_patterns_start = time.perf_counter() - config = UnifiedPatternConfig() - patterns_to_register = config.get_skeleton_patterns() - load_patterns_ms = int((time.perf_counter() - load_patterns_start) * 1000) - - if not patterns_to_register: - logger.warning("No patterns available in config") - _log_timing( - "pattern_extractor.pattern_matcher", - model=self._model.model_path, - failed=True, - reason="no_patterns_in_config", - get_model_ms=get_model_ms, - matcher_init_ms=matcher_init_ms, - load_patterns_ms=load_patterns_ms, - total_ms=int((time.perf_counter() - total_start) * 1000), - ) - return [] + def _resolve_pattern_rule_table( + self, + *, + pattern_class: str, + ep_name: str, + device: str, + preferred_domain: str, + target_opset: int, + for_debug: bool, + ) -> tuple[Path | None, str | None, int | None]: + """Resolve the most suitable parquet table for one pattern candidate.""" + search_dirs: list[Path] = [] + if for_debug: + search_dirs.extend(get_runtime_rules_debug_search_dirs()) + search_dirs.extend(get_runtime_rules_search_dirs()) + + # Keep first-seen order and skip non-existing directories. + dedup_dirs: list[Path] = [] + seen_dirs: set[Path] = set() + for base_dir in search_dirs: + try: + resolved_dir = base_dir.resolve(strict=False) + except OSError: + continue + if resolved_dir in seen_dirs or not resolved_dir.is_dir(): + continue + seen_dirs.add(resolved_dir) + dedup_dirs.append(resolved_dir) - register_start = time.perf_counter() - for pattern in patterns_to_register: - matcher.register_pattern(pattern) - register_ms = int((time.perf_counter() - register_start) * 1000) + if not dedup_dirs: + return None, None, None - logger.info("Registered %d patterns for matching", len(patterns_to_register)) + rule_subdir = f"{ep_name}_{device.upper()}" + glob_pattern = f"{pattern_class}_{ep_name}_{device.upper()}_*_opset*.parquet" - # Perform pattern matching - logger.info("Calling PatternMatcher.match()...") - match_start = time.perf_counter() - pattern_matches = matcher.match() - match_ms = int((time.perf_counter() - match_start) * 1000) - logger.info("PatternMatcher found %d matches", len(pattern_matches)) + for base_dir in dedup_dirs: + target_dir = base_dir / rule_subdir + if not target_dir.is_dir(): + continue - if not pattern_matches: - logger.info("No pattern matches found by PatternMatcher") - # Debug: try skeleton matching without validation - skeleton_results = matcher.match_skeleton() - logger.info( - "Skeleton matching found %d potential matches (before validation)", - len(skeleton_results), - ) - if skeleton_results: - matched_node_keys = skeleton_results[0].matched_node_keys - sample_nodes = matched_node_keys[:3] if matched_node_keys else [] - logger.info( - "Sample skeleton match - Pattern: %s, Nodes: %s", - skeleton_results[0].pattern.__class__.__name__, - sample_nodes, + candidates: list[tuple[Path, str, int]] = [] + for path in target_dir.glob(glob_pattern): + parsed = self._parse_pattern_rule_filename( + path.name, + pattern_class=pattern_class, + ep_name=ep_name, + device=device, ) + if parsed is None: + continue + domain_name, opset_version = parsed + candidates.append((path, domain_name, opset_version)) - logger.info( - "Extracted %d subgraph patterns using PatternMatcher", - len(pattern_matches), - ) - _log_timing( - "pattern_extractor.pattern_matcher", - model=self._model.model_path, - patterns_registered=len(patterns_to_register), - matches=len(pattern_matches), - get_model_ms=get_model_ms, - matcher_init_ms=matcher_init_ms, - load_patterns_ms=load_patterns_ms, - register_ms=register_ms, - match_ms=match_ms, - total_ms=int((time.perf_counter() - total_start) * 1000), - ) - return pattern_matches + if not candidates: + continue - def _validate_pattern_for_matching(self, pattern: SubgraphPattern) -> bool: - """Validate if pattern has required attributes for matching. + # Prefer exact-domain rows; then closest opset not above target. + same_domain_le = [c for c in candidates if c[1] == preferred_domain and c[2] <= target_opset] + if same_domain_le: + picked = max(same_domain_le, key=lambda c: c[2]) + return picked + + any_domain_le = [c for c in candidates if c[2] <= target_opset] + if any_domain_le: + picked = max(any_domain_le, key=lambda c: c[2]) + return picked + + same_domain_gt = [c for c in candidates if c[1] == preferred_domain and c[2] > target_opset] + if same_domain_gt: + picked = min(same_domain_gt, key=lambda c: c[2]) + return picked + + picked = min(candidates, key=lambda c: c[2]) + return picked + + return None, None, None + + @staticmethod + def _normalize_compile_run_cell(value: Any) -> tuple[bool, bool] | None: + """Normalize one `compile_run_success` cell to `(compile, run)` booleans.""" + raw_value = value + if not isinstance(raw_value, (list, tuple)) and hasattr(raw_value, "tolist"): + try: + raw_value = raw_value.tolist() + except Exception: # noqa: BLE001 + return None + + if not isinstance(raw_value, (list, tuple)) or len(raw_value) < 2: + return None - Args: - pattern: SubgraphPattern definition + return bool(raw_value[0]), bool(raw_value[1]) + + @staticmethod + def _extract_rule_condition_columns(column_names: list[str]) -> list[str]: + """Return parquet condition columns (excluding output metadata columns).""" + output_cols = { + "row_index", + "compile_run_success", + "compile_reason", + "run_reason", + "rule_row_count", + "case_indices", + } + return [col for col in column_names if col not in output_cols] - Returns: - True if pattern is valid for matching, False otherwise - """ - if not pattern.semantic_label: - logger.debug( - "Pattern %s has no semantic_label, skipping matching", - pattern.pattern_id, - ) - return False - return True + @staticmethod + def _normalize_case_indices(case_indices: Any) -> list[Any] | None: + """Normalize case_indices to list form for debug payloads.""" + if case_indices is None: + return None + + normalized = case_indices + if hasattr(normalized, "tolist"): + try: + normalized = normalized.tolist() + except Exception: # noqa: BLE001 + normalized = case_indices - def _create_pattern_matches( + if isinstance(normalized, list): + return normalized + if isinstance(normalized, tuple): + return list(normalized) + return [normalized] + + def _load_pattern_rule_table( self, - pattern: SubgraphPattern, - grouped_nodes: dict[str, list[tuple[str, str]]], - source_type: str, - ) -> list[PatternMatchResult]: - """Create PatternMatchResult instances from grouped nodes. + parquet_path: Path, + table_cache: dict[str, Any], + ) -> tuple[str, Any | None]: + """Load + sanitize parquet table with a per-summary cache.""" + cache_key = str(parquet_path.resolve(strict=False)).casefold() + if cache_key in table_cache: + return "ok", table_cache[cache_key] - Args: - pattern: SubgraphPattern definition - grouped_nodes: Dict mapping tag to list of (node identifier, tag) tuples - source_type: Source of the match ("hierarchy_tag" or "htp_metadata") + try: + import pandas as pd + except Exception: # noqa: BLE001 + return "pandas_unavailable", None - Returns: - List of PatternMatch instances - """ - from ...pattern.match import PatternMatchResult, SkeletonMatchResult + try: + table_df = pd.read_parquet(parquet_path) + except Exception: # noqa: BLE001 + logger.debug("Failed to read pattern parquet: %s", parquet_path, exc_info=True) + return "read_error", None - # Note: For hierarchy_tag and HTP metadata matches, we create a simplified - # PatternMatchResult without full skeleton information since these matches - # are based on tags rather than topology matching. + table_df = table_df.where(table_df.notna(), None) + for col in table_df.columns: + raw = table_df[col].to_numpy() + table_df[col] = [make_hashable(v) for v in raw] - detected_matches: list[PatternMatchResult] = [] + table_cache[cache_key] = table_df + return "ok", table_df - for tag, node_list in grouped_nodes.items(): - logger.debug( - "Found %d nodes with tag '%s' containing pattern_label '%s'", - len(node_list), - tag, - pattern.semantic_label, - ) + def _probe_candidate_pattern_mismatch( + self, + *, + candidate_pattern_obj: Any | None, + pattern_match: PatternMatchResult, + model_opsets: dict[ONNXDomain, int], + ) -> tuple[bool, str | None]: + """Probe candidate pattern preconditions via get_internal_constants_and_attributes. + + If a pattern explicitly raises PatternMismatchedError for this match, + we stop before parquet lookup and surface the mismatch reason directly. + """ + if candidate_pattern_obj is None: + return False, None - # Resolve identifiers to NodeProto and normalize to stable keys. - matched_node_identifiers = [node_identifier for node_identifier, _ in node_list] - matched_nodes = [] - matched_node_keys = [] - for node_identifier in matched_node_identifiers: - node_proto = self._model.get_node_by_key(node_identifier) - if node_proto is None: - node_proto = self._model.get_node_by_name(node_identifier) - if node_proto is None: - continue - matched_nodes.append(node_proto) - matched_node_keys.append(self._model.get_node_key(node_proto)) - - # Create a minimal SkeletonMatchResult for API compatibility - # This is a placeholder since hierarchy_tag matches don't have full skeleton info - skeleton_result = SkeletonMatchResult( - pattern=pattern, # Use the SubgraphPattern directly - matched_nodes=matched_nodes, - matched_node_keys=matched_node_keys, - matcher=None, # type: ignore - inputs=[], - output="", - removable=False, - ) + try: + schema = candidate_pattern_obj.get_schema() + except Exception: # noqa: BLE001 + return False, None - # Create PatternMatchResult with source metadata - attributes = {"source": source_type} - if source_type == "htp_metadata": - attributes["traced_tag"] = tag - else: - attributes["hierarchy_tag"] = tag - - pattern_match = PatternMatchResult( - skeleton_match_result=skeleton_result, - schema_input_to_value={}, - schema_output_to_value={}, - type_param_to_type={}, - attributes=attributes, - input_infos={}, - ) - detected_matches.append(pattern_match) + inputs: dict[str, np.ndarray] = {} + is_constant_map: dict[str, bool] = {} - return detected_matches + for input_param in schema.inputs: + input_name = input_param.name + info = pattern_match.input_infos.get(input_name) - def _match_subgraph_pattern_from_model_tags( - self, pattern: SubgraphPattern - ) -> list[PatternMatchResult]: - """Match a subgraph pattern against the model graph using hierarchy tags. + # Missing/unknown input facts means probe is inconclusive. + if info is None: + return False, None - Args: - pattern: SubgraphPattern definition + is_constant_map[input_name] = info.is_constant - Returns: - List of PatternMatchResult instances for detected matches + if info.value is not None: + inputs[input_name] = info.value + continue - Note: - This implementation matches patterns based on hierarchy_tag attributes - embedded in ONNX nodes. For nodes with hierarchy tags containing the - pattern's semantic_label, it groups them by hierarchy_tag and creates - PatternMatch instances. - """ - # Validate pattern - if not self._validate_pattern_for_matching(pattern): - return [] + if info.shape is None: + return False, None - pattern_label = pattern.semantic_label - assert pattern_label is not None # ensured by _validate_pattern_for_matching + safe_shape = tuple( + int(dim) if isinstance(dim, (int, np.integer)) and int(dim) > 0 else 1 + for dim in info.shape + ) + inputs[input_name] = np.zeros(safe_shape, dtype=np.float32) - # Get ONNX model - model_proto = self._model.get_model() - graph = model_proto.graph + try: + candidate_pattern_obj.get_internal_constants_and_attributes( + inputs=inputs, + attributes=pattern_match.attributes, + is_constant_map=is_constant_map, + domain_versions=model_opsets, + ) + except PatternMismatchedError as mismatch_error: + return True, str(mismatch_error) + except Exception: # noqa: BLE001 + logger.debug( + "Candidate mismatch probe failed for %s; continue parquet lookup", + candidate_pattern_obj.__class__.__name__, + exc_info=True, + ) - # Group nodes by hierarchy_tag that contains pattern_label - grouped_nodes: dict[str, list[tuple[str, str]]] = {} + return False, None - for node in graph.node: - # Extract hierarchy_tag attribute - hierarchy_tag = self._extract_hierarchy_tag(node) - if not hierarchy_tag: + def _query_pattern_rule_compile_run_for_match( + self, + *, + parquet_path: Path, + pattern_match: PatternMatchResult, + candidate_pattern_name: str, + model_opsets: dict[ONNXDomain, int], + table_cache: dict[str, Any], + opset_signature: tuple[tuple[str, int], ...], + query_lookup_cache: dict[ + tuple[str, str, tuple[tuple[str, Any], ...]], + tuple[ + str, + bool | None, + bool | None, + int, + int, + int, + list[Any] | None, + int, + list[str], + dict[str, Any] | None, + ], + ], + ) -> tuple[ + str, + bool | None, + bool | None, + int, + int, + int, + list[Any] | None, + int, + list[str], + dict[str, Any] | None, + ]: + """Query one candidate parquet table using one match's constraints.""" + from .runtime_checker_query import get_query_conditions_for_pattern, query_table_exact_match + + load_status, table_df = self._load_pattern_rule_table(parquet_path, table_cache) + if load_status != "ok": + return load_status, None, None, 0, 0, 0, None, 0, [], None + if table_df is None: + return "read_error", None, None, 0, 0, 0, None, 0, [], None + + row_count = int(len(table_df)) + if row_count == 0: + return "empty_table", None, None, 0, 0, 0, None, 0, [], None + + if "compile_run_success" not in table_df.columns: + return "missing_compile_run_success", None, None, row_count, 0, 0, None, 0, [], None + + match_identity = "|".join(str(key) for key in pattern_match.matched_node_keys) + if not match_identity: + match_identity = str(getattr(pattern_match, "match_id", "")) + + condition_build_cache_key = ( + match_identity, + candidate_pattern_name, + opset_signature, + ) + cached_conditions = self._query_condition_build_cache.get(condition_build_cache_key) + try: + if cached_conditions is None: + conditions, infinite_properties = get_query_conditions_for_pattern( + pattern_match=pattern_match, + pattern_name=candidate_pattern_name, + opset_versions=model_opsets, + ) + self._query_condition_build_cache[condition_build_cache_key] = ( + conditions, + infinite_properties, + ) + else: + conditions, infinite_properties = cached_conditions + except Exception: # noqa: BLE001 + logger.debug( + "Failed to build query conditions for pattern '%s'", + candidate_pattern_name, + exc_info=True, + ) + return "query_build_error", None, None, row_count, 0, 0, None, 0, [], None + + condition_columns = self._extract_rule_condition_columns(list(table_df.columns)) + query_conditions: dict[str, Any] = {} + for col in condition_columns: + if col in infinite_properties: continue + if col not in conditions: + return ( + "query_key_missing", + None, + None, + row_count, + 0, + 0, + None, + len(query_conditions), + sorted(query_conditions.keys()), + None, + ) + + encoded_value = encode_rule_condition_value_for_parquet(conditions[col]) + query_conditions[col] = make_hashable(encoded_value) - # Check if hierarchy_tag contains pattern_label - if pattern_label in hierarchy_tag: - if hierarchy_tag not in grouped_nodes: - grouped_nodes[hierarchy_tag] = [] - grouped_nodes[hierarchy_tag].append((self._model.get_node_key(node), hierarchy_tag)) - - # Create PatternMatch instances - detected_matches = self._create_pattern_matches( - pattern=pattern, - grouped_nodes=grouped_nodes, - source_type="hierarchy_tag", + query_lookup_cache_key = ( + str(parquet_path.resolve(strict=False)).casefold(), + candidate_pattern_name, + tuple(sorted(query_conditions.items())), ) + cached_query_result = query_lookup_cache.get(query_lookup_cache_key) + if cached_query_result is not None: + return cached_query_result + + if query_conditions: + matched_df = query_table_exact_match(table_df, query_conditions) + if matched_df.empty: + debug_steps: list[dict[str, Any]] = [] + current_df = table_df + first_zero_column: str | None = None + for col, value in query_conditions.items(): + rows_before = int(len(current_df)) + if col in current_df.columns: + current_df = current_df[current_df[col] == value] + rows_after = int(len(current_df)) + + debug_steps.append( + { + "column": col, + "value": repr(value), + "rows_before": rows_before, + "rows_after": rows_after, + } + ) + if first_zero_column is None and rows_after == 0: + first_zero_column = col + + debug_details = { + "type": "properties_not_found", + "pattern_name": candidate_pattern_name, + "table_path": str(parquet_path.resolve(strict=False)), + "table_file": parquet_path.name, + "total_rows": row_count, + "query_condition_count": len(query_conditions), + "query_conditions": { + key: repr(value) for key, value in query_conditions.items() + }, + "first_zero_column": first_zero_column, + "steps": debug_steps, + } + result = ( + "properties_not_found", + None, + None, + row_count, + 0, + 0, + None, + len(query_conditions), + sorted(query_conditions.keys()), + debug_details, + ) + query_lookup_cache[query_lookup_cache_key] = result + return result + matched_row = matched_df.iloc[0] + else: + matched_row = table_df.iloc[0] + + compile_run = self._normalize_compile_run_cell(matched_row.get("compile_run_success")) + if compile_run is None: + return ( + "invalid_compile_run_success", + None, + None, + row_count, + 0, + 0, + None, + len(query_conditions), + sorted(query_conditions.keys()), + None, + ) - logger.info( - "Pattern %s: found %d matches based on hierarchy_tag", - pattern.pattern_id, - len(detected_matches), + compile_ok, run_ok = compile_run + result = ( + "ok", + compile_ok, + run_ok, + row_count, + int(compile_ok), + int(run_ok), + self._normalize_case_indices(matched_row.get("case_indices")), + len(query_conditions), + sorted(query_conditions.keys()), + None, + ) + query_lookup_cache[query_lookup_cache_key] = result + return result + + @staticmethod + def _canonical_supported_status(value: str | None) -> str: + """Normalize support labels to canonical lowercase values.""" + if value is None: + return "unknow" + + normalized = str(value).strip().lower() + if normalized == "unknown": + return "unknow" + if normalized in {"supported", "partial", "unsupported", "unknow"}: + return normalized + return "unknow" + + @classmethod + def _supported_status_rank(cls, value: str | None) -> int: + """Rank support labels for descending preference ordering.""" + status = cls._canonical_supported_status(value) + rank_map = { + "supported": 3, + "partial": 2, + "unsupported": 1, + "unknow": 0, + } + return rank_map.get(status, 0) + + @classmethod + def _candidate_supported_status( + cls, + candidate: PatternRuleCompileRunResult | None, + ) -> str: + """Derive support status from one candidate compile/run snapshot.""" + if candidate is None: + return "unknow" + + if candidate.get("status") != "ok": + return "unknow" + + compile_ok = bool(candidate.get("compile")) + run_ok = bool(candidate.get("run")) + + if compile_ok and run_ok: + return "supported" + if (not compile_ok) and run_ok: + return "partial" + return "unsupported" + + @classmethod + def _match_supported_status_from_candidates( + cls, + *, + pattern_id: str, + candidate_results: list[PatternRuleCompileRunResult], + ) -> str: + """Derive one support status for a pattern match from candidate snapshots.""" + base_candidate = next( + ( + candidate + for candidate in candidate_results + if not bool(candidate.get("is_alternative", False)) + and str(candidate.get("pattern_id", "")) == pattern_id + ), + None, + ) + if base_candidate is not None: + return cls._candidate_supported_status(base_candidate) + + first_non_alternative = next( + ( + candidate + for candidate in candidate_results + if not bool(candidate.get("is_alternative", False)) + ), + None, + ) + if first_non_alternative is not None: + return cls._candidate_supported_status(first_non_alternative) + + if candidate_results: + return cls._candidate_supported_status(candidate_results[0]) + + return "unknow" + + @staticmethod + def _priority_sort_key(priority: Any) -> int: + """Convert alternative priority to sortable integer (smaller is better).""" + if isinstance(priority, bool): + return int(priority) + if isinstance(priority, int): + return priority + if isinstance(priority, str): + try: + return int(priority) + except ValueError: + pass + return 1_000_000 + + @staticmethod + def _derive_pattern_class_from_id(pattern_id: str) -> str: + """Fallback pattern class from pattern id suffix.""" + if "/" not in pattern_id: + return pattern_id + return pattern_id.split("/")[-1] + + @classmethod + def _find_alternative_candidate( + cls, + *, + candidate_results: list[PatternRuleCompileRunResult], + alt_pattern_id: str, + alt_pattern_class: str, + ) -> PatternRuleCompileRunResult | None: + """Find candidate snapshot for one configured alternative.""" + strict_match = next( + ( + candidate + for candidate in candidate_results + if bool(candidate.get("is_alternative", False)) + and str(candidate.get("pattern_id", "")) == alt_pattern_id + and str(candidate.get("pattern_class", "")) == alt_pattern_class + ), + None, + ) + if strict_match is not None: + return strict_match + + return next( + ( + candidate + for candidate in candidate_results + if bool(candidate.get("is_alternative", False)) + and str(candidate.get("pattern_id", "")) == alt_pattern_id + ), + None, ) - return detected_matches + @classmethod + def _select_and_filter_alternatives( + cls, + *, + alternatives_meta: list[dict[str, Any]], + candidate_results: list[PatternRuleCompileRunResult], + ) -> tuple[list[dict[str, Any]], list[PatternRuleCompileRunResult]]: + """Keep only one best alternative and drop unsupported-selected branches. + + Selection keys: + 1) supported_status rank: supported > partial > unsupported > unknow + 2) priority: smaller integer first + + After selecting the top alternative, if its status is ``unsupported``, + remove alternatives entirely for this pattern match. + """ + if not alternatives_meta: + base_candidates = [ + candidate + for candidate in candidate_results + if not bool(candidate.get("is_alternative", False)) + ] + return [], base_candidates + + ranked_alternatives: list[ + tuple[ + int, + int, + str, + str, + dict[str, Any], + PatternRuleCompileRunResult | None, + str, + ] + ] = [] + + for alternative in alternatives_meta: + alt_pattern_id = str(alternative.get("pattern_to_id", "")) + alt_pattern_class = str( + alternative.get("pattern_class") + or cls._derive_pattern_class_from_id(alt_pattern_id) + ) + matched_candidate = cls._find_alternative_candidate( + candidate_results=candidate_results, + alt_pattern_id=alt_pattern_id, + alt_pattern_class=alt_pattern_class, + ) + alt_status = cls._candidate_supported_status(matched_candidate) + + ranked_alternatives.append( + ( + cls._supported_status_rank(alt_status), + cls._priority_sort_key(alternative.get("priority")), + alt_pattern_id, + alt_pattern_class, + alternative, + matched_candidate, + alt_status, + ) + ) - def _extract_hierarchy_tag(self, node: onnx.NodeProto) -> str | None: - """Extract hierarchy_tag attribute from ONNX node. + ranked_alternatives.sort( + key=lambda item: ( + -item[0], + item[1], + item[2], + item[3], + ) + ) - Args: - node: ONNX NodeProto object + best_status = ranked_alternatives[0][6] + selected_alternatives: list[dict[str, Any]] = [] + selected_candidate: PatternRuleCompileRunResult | None = None - Returns: - Hierarchy tag string or None if not found - """ - for attr in node.attribute: - if attr.name == "hierarchy_tag": - return attr.s.decode("utf-8") if attr.s else None - return None + if best_status != "unsupported": + selected_alternatives = [ranked_alternatives[0][4]] + selected_candidate = ranked_alternatives[0][5] - def _match_subgraph_pattern_from_htp_metadata( - self, pattern: SubgraphPattern - ) -> list[PatternMatchResult]: - """Match a subgraph pattern using HTP metadata. + filtered_candidates: list[PatternRuleCompileRunResult] = [] + for candidate in candidate_results: + if not bool(candidate.get("is_alternative", False)): + filtered_candidates.append(candidate) + continue - This method extracts patterns from HTP metadata JSON by analyzing - the nodes mapping and module hierarchy. + if selected_candidate is not None and candidate is selected_candidate: + filtered_candidates.append(candidate) - Args: - pattern: SubgraphPattern definition + return selected_alternatives, filtered_candidates - Returns: - List of PatternMatchResult instances for detected matches - - Note: - Uses the 'nodes' section of HTP metadata which maps ONNX node names - to their traced_tag (hierarchy path). - """ - # Validate pattern - if not self._validate_pattern_for_matching(pattern): + def _build_merge_prep_metadata( + self, + *, + subgraph_patterns_by_source: dict[str, dict[str, list[PatternMatchResult]]], + model_signature: str, + ep: EPNameOrAlias | None, + device: str | None, + for_debug: bool, + on_pattern_query_result: Callable[[str, str], None] | None = None, + ) -> list[PatternMergePrepEntry]: + """Build alternatives + parquet compile/run snapshots for merge/dedup preparation.""" + if ep is None or device is None: return [] - # Load and validate HTP metadata - htp_metadata = self._load_and_validate_htp_metadata() - if not htp_metadata: + ep_name = str(ep) + device_name = device.upper() + if not self._is_valid_parquet_lookup_target(ep_name, device_name): + logger.info( + "Skip pattern parquet lookup for invalid EP/device pair: %s_%s", + ep_name, + device_name, + ) return [] - pattern_label = pattern.semantic_label - assert pattern_label is not None # ensured by _validate_pattern_for_matching - - # The 'nodes' section of HTP metadata maps node names to traced tags (str -> str). - nodes_mapping = cast("dict[str, str]", htp_metadata["nodes"]) - - # Group nodes by traced_tag that contains pattern_label - grouped_nodes = self._group_nodes_by_traced_tag( - nodes_mapping=nodes_mapping, - pattern_label=pattern_label, + cache_key = (model_signature, ep_name, device_name, bool(for_debug)) + cached_merge_prep = self._MERGE_PREP_CACHE.get(cache_key) + if cached_merge_prep is not None: + cloned = cast("list[PatternMergePrepEntry]", copy.deepcopy(cached_merge_prep)) + if on_pattern_query_result is not None: + for entry in cloned: + try: + on_pattern_query_result( + str(entry.get("pattern_id", "")), + str(entry.get("support_status", "unknow")), + ) + except Exception: + logger.debug("on_pattern_query_result callback failed", exc_info=True) + return cloned + + model_opsets = ONNXDomain.get_model_domain_opset_versions(self._model.get_model()) + source_configs: dict[str, UnifiedPatternConfig] = {} + entries: list[PatternMergePrepEntry] = [] + table_cache: dict[str, Any] = {} + query_lookup_cache: dict[ + tuple[str, str, tuple[tuple[str, Any], ...]], + tuple[ + str, + bool | None, + bool | None, + int, + int, + int, + list[Any] | None, + int, + list[str], + dict[str, Any] | None, + ], + ] = {} + parquet_resolution_cache: dict[ + tuple[str, str, str, str, int, bool], + tuple[Path | None, str | None, int | None], + ] = {} + opset_signature = tuple( + sorted((domain.value, int(version)) for domain, version in model_opsets.items()) ) - # Create PatternMatch instances - detected_matches = self._create_pattern_matches( - pattern=pattern, - grouped_nodes=grouped_nodes, - source_type="htp_metadata", - ) - - logger.info( - "Pattern %s: found %d matches from HTP metadata", - pattern.pattern_id, - len(detected_matches), - ) - - return detected_matches + for source, source_group in sorted(subgraph_patterns_by_source.items()): + if not source_group: + continue - def _load_and_validate_htp_metadata(self) -> HTPMetadata | None: - """Load and validate HTP metadata. + config = source_configs.get(source) + if config is None: + config = UnifiedPatternConfig(ihv_type=source) + source_configs[source] = config - Returns: - HTP metadata dict or None if invalid/unavailable - """ - try: - htp_metadata = self._load_htp_metadata() - except (FileNotFoundError, ValueError) as e: - logger.warning("Failed to load HTP metadata: %s", e) - return None + for pattern_class, matches in sorted(source_group.items()): + if not matches: + continue - if not htp_metadata or "nodes" not in htp_metadata: - logger.debug("No nodes section in HTP metadata") - return None + representative = matches[0] + pattern_obj = representative.pattern + pattern_id = pattern_obj.pattern_id + + config_alternatives = config.get_alternatives(pattern_obj) + alternatives_meta = [ + { + "pattern_to_id": alt.pattern_to_id, + "pattern_class": alt.pattern_class, + "priority": alt.priority, + "enabled": alt.enabled, + } + for alt in config_alternatives + ] + + alternative_priority_by_key: dict[tuple[str, str], int] = {} + alternative_priority_by_id: dict[str, int] = {} + for alternative in alternatives_meta: + alt_pattern_id = str(alternative.get("pattern_to_id", "")) + alt_pattern_class = str( + alternative.get("pattern_class") + or self._derive_pattern_class_from_id(alt_pattern_id) + ) + priority = self._priority_sort_key(alternative.get("priority")) + + key = (alt_pattern_id, alt_pattern_class) + previous_priority = alternative_priority_by_key.get(key) + if previous_priority is None or priority < previous_priority: + alternative_priority_by_key[key] = priority + + previous_id_priority = alternative_priority_by_id.get(alt_pattern_id) + if previous_id_priority is None or priority < previous_id_priority: + alternative_priority_by_id[alt_pattern_id] = priority + + candidate_specs: list[tuple[str, str, bool, Any | None]] = [ + (pattern_class, pattern_id, False, pattern_obj) + ] + seen_candidates: set[tuple[str, str]] = {(pattern_class, pattern_id)} + + for alt in config_alternatives: + alt_pattern_class = alt.pattern_class or alt.pattern_to_id.split("/")[-1] + alt_pattern_id = alt.pattern_to_id + dedup_key = (alt_pattern_class, alt_pattern_id) + if dedup_key in seen_candidates: + continue + seen_candidates.add(dedup_key) + + alt_pattern_obj: Any | None = None + if alt.pattern_class and alt.module: + try: + alt_pattern_obj = PatternConfig( + pattern_id=alt_pattern_id, + pattern_class=alt.pattern_class, + module=alt.module, + enabled=True, + ).load_pattern() + except Exception: # noqa: BLE001 + logger.debug( + "Failed to load alternative pattern %s from %s", + alt.pattern_class, + alt.module, + exc_info=True, + ) + + candidate_specs.append( + (alt_pattern_class, alt_pattern_id, True, alt_pattern_obj) + ) + + candidate_runtime_specs: list[tuple[str, str, bool, Any | None, str, int]] = [] + for candidate_class, candidate_id, is_alt, candidate_pattern_obj in candidate_specs: + if candidate_pattern_obj is not None: + preferred_domain, target_opset = self._domain_and_target_opset_for_pattern( + candidate_pattern_obj, + model_opsets, + ) + else: + preferred_domain = ONNXDomain.AI_ONNX.value + target_opset = model_opsets.get(ONNXDomain.AI_ONNX, 1) + + candidate_runtime_specs.append( + ( + candidate_class, + candidate_id, + is_alt, + candidate_pattern_obj, + preferred_domain, + int(target_opset), + ) + ) + + base_candidate_runtime_specs = [spec for spec in candidate_runtime_specs if not spec[2]] + alternative_runtime_specs = [spec for spec in candidate_runtime_specs if spec[2]] + + def _alternative_runtime_sort_key( + runtime_spec: tuple[str, str, bool, Any | None, str, int], + ) -> tuple[int, str, str]: + candidate_class, candidate_id, *_ = runtime_spec + priority = alternative_priority_by_key.get( + (candidate_id, candidate_class), + alternative_priority_by_id.get(candidate_id, 1_000_000), + ) + return (priority, candidate_id, candidate_class) + + ordered_candidate_runtime_specs = base_candidate_runtime_specs + sorted( + alternative_runtime_specs, + key=_alternative_runtime_sort_key, + ) - return htp_metadata + for match_index, pattern_match in enumerate(matches, start=1): + candidate_results: list[PatternRuleCompileRunResult] = [] + for ( + candidate_class, + candidate_id, + is_alt, + candidate_pattern_obj, + preferred_domain, + target_opset, + ) in ordered_candidate_runtime_specs: + is_mismatch, mismatch_error = self._probe_candidate_pattern_mismatch( + candidate_pattern_obj=candidate_pattern_obj, + pattern_match=pattern_match, + model_opsets=model_opsets, + ) + + if is_mismatch: + candidate_result: PatternRuleCompileRunResult = { + "pattern_class": candidate_class, + "pattern_id": candidate_id, + "is_alternative": is_alt, + "status": "mismatch_error", + "mismatch_error": mismatch_error, + "compile": None, + "run": None, + "row_count": 0, + "table_file": None, + "table_path": None, + "domain": None, + "opset_version": None, + "compile_true_rows": 0, + "run_true_rows": 0, + "case_indices": None, + "query_condition_count": 0, + "query_condition_keys": [], + "debug_details": None, + } + else: + resolution_cache_key = ( + candidate_class, + ep_name, + device_name, + preferred_domain, + target_opset, + bool(for_debug), + ) + + resolved = parquet_resolution_cache.get(resolution_cache_key) + if resolved is None: + resolved = self._resolve_pattern_rule_table( + pattern_class=candidate_class, + ep_name=ep_name, + device=device_name, + preferred_domain=preferred_domain, + target_opset=target_opset, + for_debug=for_debug, + ) + parquet_resolution_cache[resolution_cache_key] = resolved + + table_path, resolved_domain, resolved_opset = resolved + + if table_path is None: + candidate_result = { + "pattern_class": candidate_class, + "pattern_id": candidate_id, + "is_alternative": is_alt, + "status": "table_not_found", + "mismatch_error": None, + "compile": None, + "run": None, + "row_count": 0, + "table_file": None, + "table_path": None, + "domain": resolved_domain, + "opset_version": resolved_opset, + "compile_true_rows": 0, + "run_true_rows": 0, + "case_indices": None, + "query_condition_count": 0, + "query_condition_keys": [], + "debug_details": None, + } + else: + candidate_pattern_name = ( + candidate_pattern_obj.__class__.__name__ + if candidate_pattern_obj is not None + else candidate_class + ) + + ( + status, + compile_ok, + run_ok, + row_count, + compile_true_rows, + run_true_rows, + case_indices, + query_condition_count, + query_condition_keys, + debug_details, + ) = self._query_pattern_rule_compile_run_for_match( + parquet_path=table_path, + pattern_match=pattern_match, + candidate_pattern_name=candidate_pattern_name, + model_opsets=model_opsets, + table_cache=table_cache, + opset_signature=opset_signature, + query_lookup_cache=query_lookup_cache, + ) + + candidate_result = { + "pattern_class": candidate_class, + "pattern_id": candidate_id, + "is_alternative": is_alt, + "status": status, + "mismatch_error": None, + "compile": compile_ok, + "run": run_ok, + "row_count": row_count, + "table_file": table_path.name, + "table_path": str(table_path.resolve(strict=False)), + "domain": resolved_domain, + "opset_version": resolved_opset, + "compile_true_rows": compile_true_rows, + "run_true_rows": run_true_rows, + "case_indices": case_indices, + "query_condition_count": query_condition_count, + "query_condition_keys": query_condition_keys, + "debug_details": debug_details, + } + + candidate_results.append(candidate_result) + + # Alternatives are processed by priority, so the first + # supported alternative is already the optimal pick. + if is_alt and self._candidate_supported_status(candidate_result) == "supported": + break + + filtered_alternatives, filtered_candidates = self._select_and_filter_alternatives( + alternatives_meta=alternatives_meta, + candidate_results=candidate_results, + ) + + support_status = self._match_supported_status_from_candidates( + pattern_id=pattern_id, + candidate_results=filtered_candidates, + ) + + entries.append( + { + "source": source, + "pattern_class": pattern_class, + "pattern_id": pattern_id, + "match_count": len(matches), + "match_index": match_index, + "match_id": pattern_match.match_id, + "matched_node_keys": list(pattern_match.matched_node_keys), + "support_status": support_status, + "alternatives": filtered_alternatives, + "candidates": filtered_candidates, + } + ) + + if on_pattern_query_result is not None: + try: + on_pattern_query_result(pattern_id, support_status) + except Exception: + logger.debug("on_pattern_query_result callback failed", exc_info=True) + + self._MERGE_PREP_CACHE[cache_key] = cast( + "list[PatternMergePrepEntry]", + copy.deepcopy(entries), + ) + return entries - def _group_nodes_by_traced_tag( + def summary( self, - nodes_mapping: dict[str, str], - pattern_label: str, - ) -> dict[str, list[tuple[str, str]]]: - """Group nodes by their traced_tag that contains pattern_label. - - Args: - nodes_mapping: Dict mapping node names to traced tags - pattern_label: Pattern semantic label to match + ep: EPNameOrAlias | None = None, + device: str | None = None, + for_debug: bool = False, + on_pattern_query_start: Callable[[Mapping[str, int]], None] | None = None, + on_pattern_query_result: Callable[[str, str], None] | None = None, + ) -> PatternSummary: + """Generate comprehensive pattern analysis summary. Returns: - Dict mapping traced_tag to list of (node_name, traced_tag) tuples + PatternSummary with keys: + - summary: ModelStats (from model_summary()) + - subgraph_patterns: List[PatternMatchResult] (skeleton extraction + EP-priority dedup) """ - grouped_nodes: dict[str, list[tuple[str, str]]] = {} + logger.info("Generating pattern analysis summary") + total_start = time.perf_counter() - for node_name, traced_tag in nodes_mapping.items(): - if pattern_label in traced_tag: - if traced_tag not in grouped_nodes: - grouped_nodes[traced_tag] = [] - grouped_nodes[traced_tag].append((node_name, traced_tag)) + model_signature = self._compute_model_signature() + sources = self._resolve_sources_for_ep(ep) - return grouped_nodes + subgraph_patterns_by_source: dict[str, dict[str, list[PatternMatchResult]]] = {} + source_stats: list[PatternSourceStat] = [] - def get_subgraph_patterns(self) -> list[SubgraphPattern]: - """Get available subgraph pattern definitions. + for source in sources: + grouped_matches, stat = self._extract_skeleton_matches_for_source( + source=source, + model_signature=model_signature, + ) + subgraph_patterns_by_source[source] = grouped_matches + source_stats.append(stat) + + ( + subgraph_patterns_by_source, + subgraph_patterns, + ) = self._dedup_grouped_matches_for_ep( + subgraph_patterns_by_source=subgraph_patterns_by_source, + sources=sources, + model_signature=model_signature, + ep=ep, + ) - Returns: - List of SubgraphPattern objects with pattern definitions + # Build pattern count dict: pattern_id -> count + count_dict_start = time.perf_counter() + pattern_count_dict: dict[str, int] = {} + for pattern_match in subgraph_patterns: + pattern_id = pattern_match.pattern.pattern_id + pattern_count_dict[pattern_id] = pattern_count_dict.get(pattern_id, 0) + 1 + count_dict_ms = int((time.perf_counter() - count_dict_start) * 1000) - Note: - Patterns are loaded from UnifiedPatternConfig (HTPPatternRules section). - """ - logger.debug("Loading subgraph pattern definitions from UnifiedPatternConfig") + # Generate model summary with pattern count dict + metadata = self.model_summary(detected_pattern_count=pattern_count_dict) - # Load HTP patterns from UnifiedPatternConfig - config = UnifiedPatternConfig() - patterns = config.get_htp_patterns() + if on_pattern_query_start is not None: + try: + on_pattern_query_start(pattern_count_dict) + except Exception: + logger.debug("on_pattern_query_start callback failed", exc_info=True) - if not patterns: - logger.warning("No HTP patterns found in config, returning empty list") - return [] + _log_timing( + "pattern_extractor.summary", + model=self._model.model_path, + detected_subgraph_patterns=len(subgraph_patterns), + unique_pattern_ids=len(pattern_count_dict), + extract_subgraph_ms=sum(stat["elapsed_ms"] for stat in source_stats), + build_count_dict_ms=count_dict_ms, + total_ms=int((time.perf_counter() - total_start) * 1000), + ) - logger.debug("Loaded %d subgraph pattern definitions", len(patterns)) - return patterns + merge_prep = self._build_merge_prep_metadata( + subgraph_patterns_by_source=subgraph_patterns_by_source, + model_signature=model_signature, + ep=ep, + device=device, + for_debug=for_debug, + on_pattern_query_result=on_pattern_query_result, + ) + return { + "summary": metadata, + "subgraph_patterns": subgraph_patterns, + "subgraph_patterns_by_source": subgraph_patterns_by_source, + "source_stats": source_stats, + "merge_prep": merge_prep, + "model_signature": model_signature, + } def model_summary( self, diff --git a/src/winml/modelkit/analyze/core/runtime_checker.py b/src/winml/modelkit/analyze/core/runtime_checker.py index d8a04ea56..7e91b7d97 100644 --- a/src/winml/modelkit/analyze/core/runtime_checker.py +++ b/src/winml/modelkit/analyze/core/runtime_checker.py @@ -2,10 +2,10 @@ # Copyright (c) Microsoft Corporation. All rights reserved. # Licensed under the MIT License. # -------------------------------------------------------------------------- -"""RuntimeChecker - Check pattern support against runtime rules. +"""RuntimeChecker - Check operator support against runtime rules. -Implements FR-005 (Runtime support checking), FR-006 (Pattern matching), -FR-016-020 (Support classification). +Implements FR-005 (Runtime support checking) and FR-016-020 +(Support classification). """ from __future__ import annotations @@ -16,13 +16,6 @@ import tqdm -from ...pattern.config import UnifiedPatternConfig -from ..models.runtime_checks import ( - AlternativeType, - PatternAlternative, - PatternRuntime, - RuntimeTestResult, -) from ..utils.timing_utils import make_timing_logger from .runtime_checker_query import RuntimeCheckerQuery @@ -30,12 +23,9 @@ if TYPE_CHECKING: from collections.abc import Callable - import onnx - - from winml.modelkit.pattern.match import PatternMatchResult - from ...utils.constants import EPName from ..models.onnx_model import ONNXModel + from ..models.runtime_checks import PatternRuntime logger = logging.getLogger(__name__) _log_timing = make_timing_logger(logger) @@ -47,10 +37,10 @@ class RuntimeChecker: - """Check operator and subgraph pattern support against runtime rules. + """Check operator support against runtime rules. - High-level interface for checking operator-level and subgraph-level - support for a target Execution Provider (EP). + High-level interface for checking operator-level support for a target + Execution Provider (EP). Responsibilities: - Query runtime support via RuntimeCheckerQuery @@ -59,12 +49,10 @@ class RuntimeChecker: - Aggregate runtime check results FR-005: Runtime support checking - FR-006: Pattern matching against rule database FR-016-020: Support classification logic Attributes: - model: ONNX model to analyze (optional) - patterns: List of PatternMatch for subgraph detection (optional) + model: ONNX model to analyze ep: Target execution provider (e.g., "QNNExecutionProvider") device: Device string (e.g., "CPU" | "GPU" | "NPU") """ @@ -73,49 +61,47 @@ def __init__( self, ep: EPName, device: str, - model: ONNXModel | None = None, - patterns: list[PatternMatchResult] | None = None, - pattern_config: UnifiedPatternConfig | None = None, + model: ONNXModel, dynamic_axis_strict_mode: bool = False, + pattern_matched_node_status_by_key: dict[str, str] | None = None, ) -> None: """Initialize runtime checker. Args: ep: Target execution provider name device: Device string (e.g., "CPU" | "GPU" | "NPU") - model: ONNX model to analyze (optional) - patterns: List of PatternMatchResult for subgraph detection (optional) - pattern_config: Pattern configuration for reading alternatives. - If None, a default UnifiedPatternConfig is created. + model: ONNX model to analyze dynamic_axis_strict_mode: If False (default), maps any dynamic axes to (0,) for matching against first_axis test data. If True, preserves exact dynamic axis indices. + pattern_matched_node_status_by_key: Optional stable node-key -> + pattern status map (supported/partial/unsupported/unknow) + used when matched nodes bypass parquet lookup. Raises: - ValueError: If neither model nor patterns is provided + ValueError: If model is not provided """ - if model is None and patterns is None: - raise ValueError("At least one of 'model' or 'patterns' must be provided") + if model is None: + raise ValueError("'model' is required") if not device or not device.strip(): raise ValueError("device parameter cannot be empty") self._model = model - self._patterns = patterns self._ep: EPName = ep self._device = device - # Pattern configuration for reading alternatives from JSON - self._pattern_config = pattern_config or UnifiedPatternConfig() self._dynamic_axis_strict_mode = dynamic_axis_strict_mode + self._pattern_matched_node_status_by_key: dict[str, str] = dict( + pattern_matched_node_status_by_key or {} + ) # Lazy-initialized RuntimeCheckerQuery (cached for reuse) self._query: RuntimeCheckerQuery | None = None # Pre-compute rule-data availability once at construction time so that - # op_support() and subgraph_support() can read the cached result without - # repeated filesystem probes. + # op_support() can read the cached result without repeated filesystem probes. from ..utils.ep_utils import has_any_rule_data, has_rule_data_for_ep self._has_rule_data: bool = has_rule_data_for_ep(ep, device) @@ -151,6 +137,7 @@ def _get_query(self) -> RuntimeCheckerQuery: model_path=self._model.model_path, dynamic_axis_strict_mode=self._dynamic_axis_strict_mode, node_key_by_node_id=self._model.get_node_key_map(), + pattern_matched_node_status_by_key=self._pattern_matched_node_status_by_key, ) return self._query @@ -273,165 +260,23 @@ def op_support( return results - def subgraph_support( - self, - patterns: list[PatternMatchResult] | None = None, - run_unknown_op: bool = False, - ) -> list[PatternRuntime]: - """Check subgraph-level runtime support. - - Given detected patterns, check runtime support. - Each pattern returns result + optional replacement Information. - - Args: - patterns: List of PatternMatchResult objects to check. - If None, uses patterns from initialization. - for_debug: Whether to include runtime debug details for operator checks. - - Returns: - List[PatternRuntime]: Runtime results for each pattern with alternatives - - Raises: - ValueError: If patterns is None and RuntimeChecker was not initialized with patterns - """ - # Determine which patterns to use - if patterns is None: - if self._patterns is None: - raise ValueError( - "patterns parameter is required when RuntimeChecker " - "is not initialized with patterns" - ) - patterns = self._patterns - - logger.info( - "Checking subgraph pattern support via per-node operator aggregation for %d patterns", - len(patterns), - ) - - total_start = time.perf_counter() - query_pattern_total_ms = 0 - results: list[PatternRuntime] = [] - for pattern in patterns: - pattern_start = time.perf_counter() - pattern_runtime = self.query_pattern_support(pattern, run_unknown_op=run_unknown_op) - query_pattern_total_ms += int((time.perf_counter() - pattern_start) * 1000) - results.append(pattern_runtime) - - total_ms = int((time.perf_counter() - total_start) * 1000) - _log_timing( - "runtime_checker.subgraph_support", - ep=self._ep, - device=self._device, - patterns=len(results), - total_ms=total_ms, - query_pattern_ms=query_pattern_total_ms, - overhead_ms=total_ms - query_pattern_total_ms, - avg_query_pattern_ms=(query_pattern_total_ms // len(results) if results else 0), - ) - - return results - - def query_pattern_support( - self, - pattern: PatternMatchResult, - run_unknown_op: bool = False, - ) -> PatternRuntime: - """Evaluate a single pattern's runtime support + replacements. - - Args: - pattern: PatternMatchResult object to check - - Returns: - PatternRuntime: Runtime result with pattern_id, result, and alternatives - - Process: - 1. Check original pattern support via RuntimeCheckerQuery.run_for_subgraph - 2. Check possible replacement patterns (alternatives) - 3. For each alternative, evaluate its support status - 4. Return PatternRuntime with results and alternatives - """ - if self._model is None: - raise ValueError( - f"Cannot lookup pattern support for '{pattern.pattern.pattern_id}' " - f"without ONNX model. RuntimeChecker was initialized without model." - ) - - pattern_id = pattern.pattern.pattern_id - - # Get cached RuntimeCheckerQuery and check pattern support - query = self._get_query() - pattern_runtime = query.run_for_subgraph(pattern, run_unknown_op=run_unknown_op) - result = pattern_runtime.result - - logger.debug( - "Pattern %s: %s (compile=%s, run=%s)", - pattern_id, - result.classification.value, - result.compile, - result.run, - ) - - # Build alternatives from pattern config (JSON) - # TODO: Replace mock RuntimeTestResult with actual runtime checks - alternatives: list[PatternAlternative] = [] - pattern_config = self._pattern_config.get_pattern_config(pattern.pattern) - config_alternatives = self._pattern_config.get_alternatives(pattern.pattern) - for config_alt in config_alternatives: - alternative = PatternAlternative( - pattern_id=config_alt.pattern_to_id, - result=RuntimeTestResult( - compile=True, - run=True, - reason=config_alt.reason or f"Alternative for {pattern_id}", - ), - alternative_type=AlternativeType.EQUIVALENT, - enabled=config_alt.enabled, - details=config_alt.details, - action_items=config_alt.action_items, - ) - alternatives.append(alternative) - logger.debug( - "Added alternative %s for pattern %s", - config_alt.pattern_to_id, - pattern_id, - ) - - return PatternRuntime( - pattern_id=pattern_id, - result=result, - alternatives=alternatives, - explanation=pattern_config.explanation if pattern_config else None, - pattern_match=pattern, - ) - def summary( self, - patterns: list[PatternMatchResult] | None = None, for_debug: bool = False, run_unknown_op: bool = False, save_node_types: set[str] | None = None, on_node_result: Callable | None = None, ) -> dict[str, list[PatternRuntime]]: - """Combine operator-level & pattern-level runtime results. - - Args: - patterns: List of PatternMatchResult objects to check. - If None, uses patterns from initialization. + """Return operator-level runtime results. Returns: - Dict containing both op_support and subgraph_support results: - - op_runtime_check_result: Operator-level runtime check - results (only if model provided) - - subgraph_runtime_check_result: Subgraph pattern check - results + Dict containing operator-level runtime check results. """ logger.info("Generating runtime support summary") total_start = time.perf_counter() summary_dict: dict[str, list[PatternRuntime]] = {} op_support_ms = 0 - subgraph_support_ms = 0 - merge_ms = 0 # Get operator-level support (only if model is available) if self._model is not None: @@ -445,92 +290,15 @@ def summary( op_support_ms = int((time.perf_counter() - op_start) * 1000) summary_dict["op_runtime_check_result"] = op_results - # Get subgraph-level support - subgraph_start = time.perf_counter() - pattern_results = self.subgraph_support(patterns, run_unknown_op=run_unknown_op) - subgraph_support_ms = int((time.perf_counter() - subgraph_start) * 1000) - summary_dict["subgraph_runtime_check_result"] = pattern_results - - merge_start = time.perf_counter() - # Build stable node key -> PatternRuntime from pattern_results - node_to_pattern_runtime: dict[str, PatternRuntime] = {} - for pr in pattern_results: - if ( - (not pr.result.no_data) - and pr.pattern_match - and hasattr(pr.pattern_match, "skeleton_match_result") - ): - smr = pr.pattern_match.skeleton_match_result - if smr and smr.matched_node_keys: - for node_key in smr.matched_node_keys: - node_to_pattern_runtime[node_key] = pr - - # Override matching op_results - merged = [] - for op_r in summary_dict["op_runtime_check_result"]: - node_key = self._get_node_key(op_r) - if node_key in node_to_pattern_runtime: - # Replace with pattern-level result, keeping original pattern_match for traceability - pr = node_to_pattern_runtime[node_key] - merged.append( - PatternRuntime( - pattern_id=op_r.pattern_id, # keep original op pattern_id - result=pr.result, # use subgraph-level result - alternatives=[], # subgraph alternatives belong to the subgraph, not the op - pattern_match=op_r.pattern_match, - ) - ) - else: - merged.append(op_r) - - summary_dict["op_runtime_check_result"] = merged - merge_ms = int((time.perf_counter() - merge_start) * 1000) - total_ms = int((time.perf_counter() - total_start) * 1000) _log_timing( "runtime_checker.summary", ep=self._ep, device=self._device, op_results=len(summary_dict.get("op_runtime_check_result", [])), - subgraph_results=len(summary_dict.get("subgraph_runtime_check_result", [])), total_ms=total_ms, op_support_ms=op_support_ms, - subgraph_support_ms=subgraph_support_ms, - merge_ms=merge_ms, - overhead_ms=total_ms - op_support_ms - subgraph_support_ms - merge_ms, + overhead_ms=total_ms - op_support_ms, ) return summary_dict - - def _get_node_key(self, op_runtime: PatternRuntime) -> str: - """Extract stable node key from an op-level PatternRuntime.""" - pm = op_runtime.pattern_match - if pm and hasattr(pm, "skeleton_match_result"): - node_keys: list[str] = pm.skeleton_match_result.matched_node_keys - if node_keys: - return node_keys[0] - return "" - - def _make_op_key(self, node: onnx.NodeProto) -> str: - """Generate operator key from node. - - Internal method to create unique key for operator. - - Args: - node: ONNX node - - Returns: - Operator key string (e.g., "OP/ai.onnx/Conv") - - Note: - This is an internal method. - """ - # Detect namespace - namespace = "ai.onnx" # Default namespace - if node.domain: - if node.domain == "com.microsoft": - namespace = "com.microsoft" - elif node.domain != "": - namespace = node.domain - - return f"OP/{namespace}/{node.op_type}" diff --git a/src/winml/modelkit/analyze/core/runtime_checker_query.py b/src/winml/modelkit/analyze/core/runtime_checker_query.py index 454a67422..84f5712da 100644 --- a/src/winml/modelkit/analyze/core/runtime_checker_query.py +++ b/src/winml/modelkit/analyze/core/runtime_checker_query.py @@ -987,6 +987,7 @@ def __init__( model_path: str | Path | None = None, dynamic_axis_strict_mode: bool = False, node_key_by_node_id: dict[int, str] | None = None, + pattern_matched_node_status_by_key: dict[str, str] | None = None, ) -> None: """Initialize runtime checker query. @@ -1000,6 +1001,9 @@ def __init__( for matching against first_axis test data. If True, preserves exact dynamic axis indices. node_key_by_node_id: Optional sidecar map from id(node) to stable node key. + pattern_matched_node_status_by_key: Optional stable node-key to + pattern status mapping (supported/partial/unsupported/unknow) + used to classify matched nodes when parquet lookup is skipped. """ self.model_path = str(Path(model_path).resolve(strict=False)) if model_path else None self.model_base_dir = str(Path(self.model_path).parent) if self.model_path else None @@ -1042,6 +1046,11 @@ def __init__( else: self._node_key_by_node_id = build_node_key_by_node_id(self._graph_nodes) + self._pattern_matched_node_status_by_key: dict[str, str] = { + str(node_key): str(status) + for node_key, status in (pattern_matched_node_status_by_key or {}).items() + } + self.ep_name = ep_name self.device_type = device_type self.valueinfo = collect_valueinfo_dict(self.model_proto) @@ -1089,6 +1098,49 @@ def __init__( # since_version cache keyed by (op, domain, model_opset) self._since_version_cache: dict[tuple[str, str, int], int] = {} + @staticmethod + def _build_op_pattern_id(node: onnx.NodeProto) -> str: + """Build OP// identifier for one ONNX node.""" + try: + op_domain = ONNXDomain.from_str(node.domain) + domain_value = op_domain.value + except ValueError: + domain_value = node.domain or ONNXDomain.AI_ONNX.value + + return f"OP/{domain_value}/{node.op_type}" + + @staticmethod + def _runtime_result_from_pattern_status(pattern_status: str) -> RuntimeTestResult: + """Map pattern status string to RuntimeTestResult for matched nodes.""" + normalized = (pattern_status or "unknow").strip().lower() + if normalized == "supported": + return RuntimeTestResult( + compile=True, + run=True, + no_data=False, + reason="pattern_matched", + ) + if normalized == "partial": + return RuntimeTestResult( + compile=False, + run=True, + no_data=False, + reason="pattern_matched", + ) + if normalized == "unsupported": + return RuntimeTestResult( + compile=False, + run=False, + no_data=False, + reason="pattern_matched", + ) + return RuntimeTestResult( + compile=False, + run=False, + no_data=True, + reason="pattern_matched", + ) + def _collect_qdq_types(self) -> None: """Collect QDQ types from the model. @@ -1535,78 +1587,6 @@ def _generate_model_inputs(self, model: onnx.ModelProto) -> dict[str, np.ndarray return input_feed - def _generate_node_inputs(self, node: onnx.NodeProto) -> dict[str, np.ndarray]: - """Generate dummy input data for a single-node model. - - Creates numpy arrays with appropriate shapes and dtypes based on the - node's input value info. Initializer/constant inputs are excluded since - they are embedded in the model. - - Args: - node: The ONNX node to generate inputs for. - - Returns: - Dict mapping input names to numpy arrays. - - Raises: - ValueError: If dtype or shape information is missing for an input. - """ - input_feed: dict[str, np.ndarray] = {} - default_dim_size = 2 # Replace dynamic/unknown dims with this size - - for inp_name in node.input: - if not inp_name: - continue - # Skip regular initializers/constants - they are embedded in the model. - # External-data initializers are modeled as runtime inputs. - if inp_name in self.initializers: - init = self.initializers[inp_name] - if init.data_location != onnx.TensorProto.EXTERNAL: - continue - - try: - np_dtype = onnx.helper.tensor_dtype_to_np_dtype(init.data_type) - except Exception: - np_dtype = np.dtype(np.float32) - - shape = tuple(int(d) for d in init.dims) - input_feed[inp_name] = np.zeros(shape, dtype=np_dtype) - continue - - if inp_name in self.constants: - continue - - vi = self.valueinfo.get(inp_name) - if vi is None: - raise ValueError( - f"Input '{inp_name}' for node '{node.name}' ({node.op_type}) " - f"not found in valueinfo" - ) - - vi_shape, dtype_str = shape_and_dtype_from_valueinfo(vi) - if dtype_str is None: - raise ValueError( - f"Input '{inp_name}' for node '{node.name}' ({node.op_type}) " - f"has no dtype information" - ) - - # Convert dtype string to numpy dtype - np_dtype = SupportedONNXType.from_annotation(dtype_str).np_type - - concrete_shape: tuple[int, ...] - if vi_shape is None: - # No shape info at all - use a simple 1D array - concrete_shape = (default_dim_size,) - else: - # Replace dynamic dimensions (strings or None) with default size - concrete_shape = tuple( - d if isinstance(d, int) and d > 0 else default_dim_size for d in vi_shape - ) - - input_feed[inp_name] = np.zeros(concrete_shape, dtype=np_dtype) - - return input_feed - def _try_local_ep_check( self, node: onnx.NodeProto, @@ -1883,15 +1863,6 @@ def _save_failed_node( except Exception as e: logger.warning("Failed to save node for %s: %s", node.op_type, e) - def run_for_model_per_op(self) -> dict[str, Any]: - """Run runtime check for all nodes in model. - - Returns: - Dict with results for each operator - """ - # run run_for_nodes for all nodes - return {} - def _maybe_save_failed_node_result( self, node: onnx.NodeProto, @@ -2345,10 +2316,6 @@ def run_for_node( ), ) - pattern_match_start = time.perf_counter() - pattern_match = node_to_pattern_match(node, node_key) - pattern_match_ms = _elapsed_ms(pattern_match_start) - def _finish(result: PatternRuntime, outcome: str) -> PatternRuntime: _log_timing( "run_for_node", @@ -2373,6 +2340,36 @@ def _finish(result: PatternRuntime, outcome: str) -> PatternRuntime: ) return result + if node_key in self._pattern_matched_node_status_by_key: + pattern_status = self._pattern_matched_node_status_by_key[node_key] + pattern_matched_debug_details: RuntimeDebugDetails | None = None + if for_debug: + pattern_matched_debug_details = { + "type": "pattern_matched", + "node_stable_key": node_key, + "op_type": node.op_type, + "status": pattern_status, + "table_path": None, + "table_file": None, + } + + result = self._runtime_result_from_pattern_status(pattern_status) + result.debug_details = pattern_matched_debug_details + + return _finish( + PatternRuntime( + pattern_id=self._build_op_pattern_id(node), + result=result, + alternatives=self.alternatives, + pattern_match=None, + ), + outcome="pattern_matched", + ) + + pattern_match_start = time.perf_counter() + pattern_match = node_to_pattern_match(node, node_key) + pattern_match_ms = _elapsed_ms(pattern_match_start) + # Ignore QuantizeLinear and DequantizeLinear ops for now, # Q and DQ ops will be tested in quantized ops ignored_ops = { @@ -2560,156 +2557,3 @@ def get_pattern_id(is_qdq: bool) -> str: ) parquet_rules_ms = _elapsed_ms(parquet_rules_start) return _finish(final_result, outcome="parquet_rules") - - def run_for_subgraph( - self, - pattern_match: PatternMatchResult, - run_unknown_op: bool = False, - ) -> PatternRuntime: - """Run runtime check for subgraph pattern via per-node checks.""" - pattern_name = pattern_match.pattern.__class__.__name__ - logger.debug( - "Pattern-level aggregated rules are removed; checking individual operators for '%s'", - pattern_name, - ) - return self._run_for_subgraph_per_node( - pattern_match, - pattern_name, - run_unknown_op, - ) - - def _run_for_subgraph_per_node( - self, - pattern_match: PatternMatchResult, - pattern_name: str, - run_unknown_op: bool, - ) -> PatternRuntime: - """Fallback: check each operator in the pattern individually. - - Args: - pattern_match: PatternMatchResult containing pattern information. - pattern_name: Pattern variant name. - run_unknown_op: If True, attempt local EP check for unknown ops. - - Returns: - PatternRuntime with aggregated results from individual node checks. - """ - pattern_id = pattern_match.pattern.pattern_id - - if ( - not hasattr(pattern_match, "skeleton_match_result") - or pattern_match.skeleton_match_result is None - ): - logger.warning( - f"Pattern '{pattern_id}' has no " - f"skeleton_match_result, cannot check " - f"individual nodes" - ) - return PatternRuntime( - pattern_id=pattern_id, - result=RuntimeTestResult( - compile=False, - run=False, - no_data=True, - reason=( - f"Pattern '{pattern_name}' not " - f"found in database and has no " - f"matched nodes to check" - ), - debug_details=None, - ), - alternatives=self.alternatives, - pattern_match=pattern_match, - ) - - matched_nodes = pattern_match.skeleton_match_result.matched_nodes - - if not matched_nodes: - logger.warning("Pattern '%s' has no matched nodes", pattern_id) - return PatternRuntime( - pattern_id=pattern_id, - result=RuntimeTestResult( - compile=False, - run=False, - no_data=True, - reason=f"Pattern '{pattern_name}' has no nodes to check", - debug_details=None, - ), - alternatives=self.alternatives, - pattern_match=pattern_match, - ) - - # Check runtime support for each node in the pattern - node_results: list[PatternRuntime] = [] - for node in matched_nodes: - node_result = self.run_for_node(node, run_unknown_op=run_unknown_op) - node_results.append(node_result) - - # Aggregate results: pattern is supported only if ALL nodes are supported - all_compile = all(r.result.compile for r in node_results) - all_run = all(r.result.run for r in node_results) - any_no_data = any(r.result.no_data for r in node_results) - - # Collect failure reasons - failed_nodes = [ - f"{r.pattern_id}: {r.result.reason}" - for r in node_results - if not r.result.compile or not r.result.run - ] - - no_data_nodes = [r.pattern_id for r in node_results if r.result.no_data] - - if all_compile and all_run and not any_no_data: - return PatternRuntime( - pattern_id=pattern_id, - result=RuntimeTestResult( - compile=True, - run=True, - no_data=False, - reason=( - f"Pattern '{pattern_name}' fully " - f"supported: all " - f"{len(node_results)} operators " - f"supported" - ), - debug_details=None, - ), - alternatives=self.alternatives, - pattern_match=pattern_match, - ) - - if any_no_data: - return PatternRuntime( - pattern_id=pattern_id, - result=RuntimeTestResult( - compile=False, - run=False, - no_data=True, - reason=( - f"Pattern '{pattern_name}' status " - f"unknown: no data for operators " - f"{', '.join(no_data_nodes[:3])}" - f"{'...' if len(no_data_nodes) > 3 else ''}" - ), - debug_details=None, - ), - alternatives=self.alternatives, - pattern_match=pattern_match, - ) - - failure_summary = "; ".join(failed_nodes[:3]) - if len(failed_nodes) > 3: - failure_summary += f" (and {len(failed_nodes) - 3} more)" - - return PatternRuntime( - pattern_id=pattern_id, - result=RuntimeTestResult( - compile=all_compile, - run=all_run, - no_data=False, - reason=f"Pattern '{pattern_name}' has unsupported operators: {failure_summary}", - debug_details=None, - ), - alternatives=self.alternatives, - pattern_match=pattern_match, - ) diff --git a/src/winml/modelkit/analyze/models/runtime_checks.py b/src/winml/modelkit/analyze/models/runtime_checks.py index 3329acdad..c697e331a 100644 --- a/src/winml/modelkit/analyze/models/runtime_checks.py +++ b/src/winml/modelkit/analyze/models/runtime_checks.py @@ -32,6 +32,7 @@ class RuntimeDebugDetails(TypedDict): source: NotRequired[str] fallback_reason: NotRequired[str] op_type: NotRequired[str] + status: NotRequired[str] node_stable_key: NotRequired[str | None] domain: NotRequired[str] opset_version: NotRequired[int] diff --git a/src/winml/modelkit/analyze/runtime_checker/result_processor.py b/src/winml/modelkit/analyze/runtime_checker/result_processor.py index 23d27c5e5..e2561db9e 100644 --- a/src/winml/modelkit/analyze/runtime_checker/result_processor.py +++ b/src/winml/modelkit/analyze/runtime_checker/result_processor.py @@ -10,7 +10,7 @@ import numpy as np import pandas as pd -from onnx.defs import SchemaError, onnx_opset_version +from onnx.defs import SchemaError from ...onnx import ONNXDomain from ...pattern.base import get_pattern_input_generator @@ -21,7 +21,6 @@ ) from ..utils.model_utils import ( encode_rule_condition_value_for_parquet, - get_op_since_version, make_hashable, ) from ..utils.rule_loader import get_runtime_rules_search_dirs @@ -31,64 +30,6 @@ from ...utils.constants import EPName -# Snapshot metadata keys used in generated rule artifacts. -SNAPSHOT_TYPE_KEY = "__snapshot_type__" -SNAPSHOT_TYPE_DELTA = "delta_v1" -SNAPSHOT_BASE_OPSET_KEY = "__base_opset__" -SNAPSHOT_CURRENT_OPSET_KEY = "__current_opset__" -SNAPSHOT_CHANGED_KEY = "__changed__" -SNAPSHOT_DELETED_KEY = "__deleted__" - - -def _sorted_dict_by_key(payload: dict[str, Any]) -> dict[str, Any]: - """Return a shallow key-sorted dict for stable JSON output.""" - return dict(sorted(payload.items())) - - -def _build_snapshot_payload( - current_payload: dict[str, Any], - current_opset: int, - previous_payload: dict[str, Any] | None, - previous_opset: int | None, -) -> dict[str, Any]: - """Build either a full snapshot (first version) or a delta snapshot. - - Full snapshots keep backward compatibility with existing plain-dict format. - Delta snapshots store only changed/deleted operators relative to the previous opset. - """ - if previous_payload is None or previous_opset is None: - return _sorted_dict_by_key(current_payload) - - changed = { - op_name: value - for op_name, value in current_payload.items() - if op_name not in previous_payload or previous_payload[op_name] != value - } - deleted = sorted(op_name for op_name in previous_payload if op_name not in current_payload) - - return { - SNAPSHOT_TYPE_KEY: SNAPSHOT_TYPE_DELTA, - SNAPSHOT_BASE_OPSET_KEY: previous_opset, - SNAPSHOT_CURRENT_OPSET_KEY: current_opset, - SNAPSHOT_CHANGED_KEY: _sorted_dict_by_key(changed), - SNAPSHOT_DELETED_KEY: deleted, - } - - -def _is_delta_snapshot_payload(payload: Any) -> bool: - return isinstance(payload, dict) and payload.get(SNAPSHOT_TYPE_KEY) == SNAPSHOT_TYPE_DELTA - - -def _can_append_merge(existing_payload: Any, new_payload: Any) -> bool: - """Whether append-mode shallow dict merge is safe for these payloads.""" - return ( - isinstance(existing_payload, dict) - and isinstance(new_payload, dict) - and not _is_delta_snapshot_payload(existing_payload) - and not _is_delta_snapshot_payload(new_payload) - ) - - def _get_input_constraint_types( check_results: list[dict[str, Any]], ) -> dict[str, str]: @@ -442,101 +383,6 @@ def extract_single_negative_rules( return all_negative_rules, all_failed -def build_op_query_negative_rules_and_table( - check_results: list[dict[str, Any]], - input_generator: OpInputGenerator, - use_qdq: bool, - op_version: int, - device: str, - ep_name: EPName, - op_domain: str, - # schema: OpSchema, -) -> tuple[dict[str, Any], pd.DataFrame]: - """Build negative rules from check results for a specific operator. - - Args: - check_results: List of check result items from runtime checker - input_generator: OpInputGenerator object for the operator - - Returns: - Tuple of (negative_rules_dict, dataframe): - - negative_rules_dict: Dictionary containing operator name and negative rules - - dataframe: DataFrame with all test results and properties - """ - op_name = input_generator.op_name - if not check_results: - return {"op_name": op_name, "negative_rules": {}}, pd.DataFrame() - - # Convert items to rows - - # Pre-compute constraint types from non-None constraints for consistent property naming - input_constraint_types = _get_input_constraint_types(check_results) - # Pre-compute all attribute names for consistent property naming - all_attr_names = _get_all_attr_names(check_results) - - def get_row(item: dict[str, Any]) -> dict[str, Any]: - """Convert item to row with derived properties if available.""" - row = item_to_row( - item, - input_constraint_types, - all_attr_names, - input_generator.replace_float_with_dummy_in_query, - use_qdq=use_qdq, - ) - try: - row = input_generator.derive_properties(row) - except NotImplementedError: - pass - return row - - rows = [get_row(item) for item in check_results] - - # Create DataFrame and replace NaN with None - df = pd.DataFrame(rows, dtype=object) - df = df.replace({np.nan: None}) - - # Auto-detect infinite properties (those ending with _shape or _value) - # These represent unbounded input spaces that should not be used for negative rules - infinite_properties = input_generator.get_infinite_property_names() - internal_reason_cols = [ - "compile_reason", - "run_reason", - "has_not_run_placeholder_reason", - "case_index", - ] - consistency_ignored = [*infinite_properties, *internal_reason_cols] - assert check_df_consistent( - df, - op_name, - "compile_run_success", - consistency_ignored, - op_version=op_version, - device=device, - ep_name=ep_name, - op_domain=op_domain, - is_qdq=use_qdq, - ) - - # Internal reason columns are only for consistency filtering and must not be - # exported to tables/rules, otherwise downstream matcher treats them as - # required condition keys. - export_df = df.drop(columns=internal_reason_cols, errors="ignore") - - negative_rules, all_failed = extract_single_negative_rules( - export_df, "compile_run_success", infinite_properties - ) - names = ["compile", "run"] - - negative_rules_dict = { - "op_name": op_name, - "negative_rules": dict(zip(names, negative_rules, strict=False)), - "all_failed": dict(zip(names, all_failed, strict=False)), - "total_row_count": len(export_df), - } - - return negative_rules_dict, export_df - - def _parse_filename(filename: str) -> tuple[str, str, str, str, int, bool]: """Parse operator name, EP name, domain, opset, and QDQ flag from filename. @@ -579,39 +425,6 @@ def _parse_filename(filename: str) -> tuple[str, str, str, str, int, bool]: return op_domain, op_name, ep_name, device, opset_version, is_qdq -def get_opset_version_range(op_name: str, start_opset_version: int, op_domain: str) -> list[int]: - """Get the range of opset versions that use the same op schema version. - - Given an op_name and a starting opset version, determines all consecutive opset - versions that use the same since_version of the operator. This is useful when - updating rules: e.g., if Slice has versions 1, 10, 11, 13, and start_opset_version=11, - the since_version is 11 and the next version is 13, so we return [11, 12]. - - Args: - op_name: Name of the ONNX operator (e.g., "Slice") - start_opset_version: The starting opset version - op_domain: The domain of the operator (empty string for ai.onnx) - - Returns: - List of consecutive opset versions sharing the same op schema version - """ - max_opset = onnx_opset_version() - base_since = get_op_since_version(op_name, start_opset_version, op_domain) - - versions = [] - for v in range(start_opset_version, max_opset + 1): - try: - since = get_op_since_version(op_name, v, op_domain) - except SchemaError: - break - if since == base_since: - versions.append(v) - else: - break - - return versions - - def _parse_requested_domains(domains_arg: str) -> list[str]: """Parse and validate --domains values.""" requested_domains = [part.strip() for part in domains_arg.split(",") if part.strip()] diff --git a/src/winml/modelkit/analyze/utils/json_utils.py b/src/winml/modelkit/analyze/utils/json_utils.py index 8acc40c9b..b7f79c5a3 100644 --- a/src/winml/modelkit/analyze/utils/json_utils.py +++ b/src/winml/modelkit/analyze/utils/json_utils.py @@ -32,35 +32,3 @@ def validate_json_schema(data: dict[str, Any], schema_path: Path) -> bool: validate(instance=data, schema=schema) return True - - -def load_json_file(file_path: Path) -> dict[str, Any]: - """Load and parse JSON file. - - Args: - file_path: Path to JSON file - - Returns: - Parsed JSON data as dictionary - - Raises: - FileNotFoundError: If file not found - json.JSONDecodeError: If JSON is malformed - """ - if not file_path.exists(): - raise FileNotFoundError(f"JSON file not found: {file_path}") - - return json.loads(file_path.read_text(encoding="utf-8")) # type: ignore[no-any-return] - - -def save_json_file(data: dict[str, Any], file_path: Path, indent: int = 2) -> None: - """Save data to JSON file. - - Args: - data: Data to serialize - file_path: Output file path - indent: JSON indentation level - """ - file_path.parent.mkdir(parents=True, exist_ok=True) - - file_path.write_text(json.dumps(data, indent=indent, ensure_ascii=False), encoding="utf-8") diff --git a/src/winml/modelkit/analyze/utils/op_utils.py b/src/winml/modelkit/analyze/utils/op_utils.py index 4b9d4cf03..6b94458fb 100644 --- a/src/winml/modelkit/analyze/utils/op_utils.py +++ b/src/winml/modelkit/analyze/utils/op_utils.py @@ -146,12 +146,6 @@ def _compute_case_index_with_namespace_key(case: dict, *, namespace_key: str) -> return f"{namespace_key}{_hash_case_signature(signature)}" -def compute_case_index(case: dict, *, namespace: str) -> str: - """Compute unified 36-char case_index for a case under the given file namespace.""" - namespace_key = encode_file_name_to_4char_key(namespace) - return _compute_case_index_with_namespace_key(case, namespace_key=namespace_key) - - class CheckResultWriter: """Writer for test results that supports continuation from existing files.""" @@ -230,10 +224,6 @@ def __init__( if not (compile_success and run_success): self.failed_signatures.add(sig) - def has_existing_results(self) -> bool: - """Check if we have existing results to work with.""" - return len(self.existing_signatures) > 0 - def should_skip_case(self, case: dict) -> bool: """Check if a case should be skipped based on its signature. diff --git a/src/winml/modelkit/analyze/utils/rule_loader.py b/src/winml/modelkit/analyze/utils/rule_loader.py index 1729605d0..d0dbea7a3 100644 --- a/src/winml/modelkit/analyze/utils/rule_loader.py +++ b/src/winml/modelkit/analyze/utils/rule_loader.py @@ -154,14 +154,6 @@ def __init__(self, rules_dir: Path | None = None) -> None: self.rules_dir = Path(rules_dir) self.runtime_rules: dict[str, list[RuntimeCheckRule]] = {} - def get_runtime_rules_dir(self) -> Path: - """Get the path to runtime check rules directory. - - Returns: - Path to runtime_check_rules directory - """ - return self.rules_dir / "runtime_check_rules" - def load_runtime_rules( self, ihv_type: IHVType | None = None ) -> dict[str, list[RuntimeCheckRule]]: diff --git a/src/winml/modelkit/cli.py b/src/winml/modelkit/cli.py index fe36e905c..5ecf696b0 100644 --- a/src/winml/modelkit/cli.py +++ b/src/winml/modelkit/cli.py @@ -305,4 +305,4 @@ def _shutdown_telemetry() -> None: if __name__ == "__main__": - main() + main() # type: ignore diff --git a/src/winml/modelkit/commands/analyze.py b/src/winml/modelkit/commands/analyze.py index dfa18bb6f..bdaf0758d 100644 --- a/src/winml/modelkit/commands/analyze.py +++ b/src/winml/modelkit/commands/analyze.py @@ -17,6 +17,7 @@ import logging import os import re +import time from pathlib import Path from typing import TYPE_CHECKING, Any, Literal, cast @@ -71,6 +72,7 @@ _TRAILING_PAREN_RE = re.compile(r" \([^()]*\)$") _RUNTIME_DEBUG_LEVELS = ("unsupported", "partial", "supported") +_SUPPORT_LEVEL_KEYS = ("supported", "partial", "unsupported", "unknown") def _display_name(pattern_id: str) -> str: @@ -106,19 +108,19 @@ def _worst_level_icon(counts: dict[str, int]) -> str: def _build_stacked_bar(counts: dict[str, int], max_count: int) -> Text: """Build a stacked bar where total width is proportional to max_count.""" - total = sum(counts.values()) + total = sum(counts.get(level, 0) for level in _SUPPORT_LEVEL_KEYS) if total == 0: return Text() bar_width = max(1, round(total / max_count * MAX_BAR_WIDTH)) # Ensure bar can fit all non-zero segments - nonzero = sum(1 for v in counts.values() if v > 0) + nonzero = sum(1 for level in _SUPPORT_LEVEL_KEYS if counts.get(level, 0) > 0) bar_width = max(bar_width, nonzero) bar = Text() chars_used = 0 - for level in ("supported", "partial", "unsupported", "unknown"): + for level in _SUPPORT_LEVEL_KEYS: count = counts.get(level, 0) if count == 0: continue @@ -130,24 +132,73 @@ def _build_stacked_bar(counts: dict[str, int], max_count: int) -> Text: return bar -def _build_analyzed_text(counts: dict[str, int]) -> Text: - """Build 'S/P/U/Unk' format like '53/0/0/0' or '12/5/1/3' with colors.""" - w = counts.get("supported", 0) - g = counts.get("partial", 0) - b = counts.get("unsupported", 0) - u = counts.get("unknown", 0) +def _build_support_text(counts: dict[str, int]) -> Text: + """Build 'S/P/U/Unk' format with per-level colors.""" + supported_count = counts.get("supported", 0) + partial_count = counts.get("partial", 0) + unsupported_count = counts.get("unsupported", 0) + unknown_count = counts.get("unknown", 0) text = Text() - text.append(str(w), style="bold green") + text.append(str(supported_count), style="bold green") text.append("/", style="dim") - text.append(str(g), style="bold yellow" if g > 0 else "dim") + text.append(str(partial_count), style="bold yellow" if partial_count > 0 else "dim") text.append("/", style="dim") - text.append(str(b), style="bold red" if b > 0 else "dim") + text.append(str(unsupported_count), style="bold red" if unsupported_count > 0 else "dim") text.append("/", style="dim") - text.append(str(u), style="bold bright_black" if u > 0 else "dim") + text.append(str(unknown_count), style="bold bright_black" if unknown_count > 0 else "dim") return text +def _format_count_breakdown( + *, + counts_by_item: dict[str, int], + max_items: int = 8, +) -> str: + """Build compact breakdown text like A(1)+B(2)+...""" + ranked_items = sorted( + ((name, int(count)) for name, count in counts_by_item.items() if int(count) > 0), + key=lambda item: (-item[1], item[0]), + ) + if not ranked_items: + return "" + + displayed_items = ranked_items[:max_items] + tokens = [ + f"{name.split('/')[-1]}({count})" + for name, count in displayed_items + ] + if len(ranked_items) > max_items: + tokens.append("...") + + return "+".join(tokens) + + +def _build_pattern_coverage_op_line(ep_payload: dict[str, Any]) -> str: + """Build one-line internal-op coverage summary for PATTERN CHECK.""" + op_counts: dict[str, int] = {} + + pattern_items = ep_payload.get("patterns", []) if isinstance(ep_payload, dict) else [] + for pattern_item in pattern_items: + node_breakdown = pattern_item.get("node_breakdown", []) + if not isinstance(node_breakdown, list): + continue + for breakdown_item in node_breakdown: + if not isinstance(breakdown_item, dict): + continue + op_type = str(breakdown_item.get("op_type", "")).strip() + total_count = int(breakdown_item.get("total_count", 0)) + if not op_type or total_count <= 0: + continue + op_counts[op_type] = op_counts.get(op_type, 0) + total_count + + breakdown = _format_count_breakdown(counts_by_item=op_counts) + total_op_count = sum(op_counts.values()) + if not breakdown: + return "Coverage OP(0)=(none)" + return f"Coverage OP({total_op_count})={breakdown}" + + def _build_analysis_table( data: dict[str, dict[str, int]], ep_device_pair_display_name: str | None = None, @@ -198,7 +249,10 @@ def _build_analysis_table( if all_ops: max_count = max(all_ops.values(), default=1) else: - max_count = max((sum(v.values()) for v in data.values()), default=1) + max_count = max( + (sum(v.get(level, 0) for level in _SUPPORT_LEVEL_KEYS) for v in data.values()), + default=1, + ) table = Table( title=title, @@ -210,10 +264,15 @@ def _build_analysis_table( ) table.add_column("Op Type", width=28, no_wrap=True) - table.add_column("S/P/U/Unk", width=16, no_wrap=True) - table.add_column("", no_wrap=True) - - agg: dict[str, int] = {"supported": 0, "partial": 0, "unsupported": 0, "unknown": 0} + table.add_column("S/P/U/Unk", width=20, no_wrap=True) + table.add_column("", no_wrap=False) + + agg: dict[str, int] = { + "supported": 0, + "partial": 0, + "unsupported": 0, + "unknown": 0, + } for op_type in display_order: total = all_ops.get(op_type, 0) if all_ops else sum(data.get(op_type, {}).values()) @@ -229,7 +288,7 @@ def _build_analysis_table( ) else: # Has data — show progress (partial or complete) - analyzed_for_op = sum(counts.values()) + analyzed_for_op = sum(counts.get(level, 0) for level in _SUPPORT_LEVEL_KEYS) for level in agg: agg[level] += counts.get(level, 0) @@ -249,12 +308,16 @@ def _build_analysis_table( remaining_width = max(1, round(remaining / max_count * MAX_BAR_WIDTH)) bar.append("░" * remaining_width, style="dim") - table.add_row(op_label, _build_analyzed_text(counts), bar) + table.add_row(op_label, _build_support_text(counts), bar) # Summary row table.add_section() - total_ops = sum(all_ops.values()) if all_ops else sum(agg.values()) - analyzed_count = sum(agg.values()) + total_ops = ( + sum(all_ops.values()) + if all_ops + else sum(agg.get(level, 0) for level in _SUPPORT_LEVEL_KEYS) + ) + analyzed_count = sum(agg.get(level, 0) for level in _SUPPORT_LEVEL_KEYS) total_label = Text() total_label.append("TOTAL", style="bold") if analyzed_count < total_ops: @@ -271,97 +334,164 @@ def _build_analysis_table( table.add_row( total_label, - _build_analyzed_text(agg), + _build_support_text(agg), total_bar, ) return table -_STATUS_ICONS = {"s": "🟢", "p": "🟡", "u": "🔴", "uk": "🔵"} -_PATTERN_STATUS_LABELS = {"s": "supported", "p": "partial", "u": "unsupported", "uk": "unknown"} -_SUPPORT_LEVEL_TO_SHORT = { - "supported": "s", - "partial": "p", - "unsupported": "u", - "unknown": "uk", -} +def _build_pattern_query_table( + data: dict[str, dict[str, int]], + ep_device_pair_display_name: str | None = None, + complete: bool = False, + all_patterns: dict[str, int] | None = None, +) -> Table: + """Build pattern query progress table with S/P/U/Unk counts.""" + title = "📊 PATTERN CHECK" + if ep_device_pair_display_name: + title += f" — [bold cyan]{ep_device_pair_display_name}[/bold cyan]" + if complete: + title += " [bold green]✅ Complete[/bold green]" + if all_patterns: + display_order = sorted(all_patterns, key=lambda x: all_patterns[x], reverse=True) + max_count = max(all_patterns.values(), default=1) + else: + display_order = sorted(data, key=lambda x: sum(data[x].values()), reverse=True) + max_count = max( + (sum(v.get(level, 0) for level in _SUPPORT_LEVEL_KEYS) for v in data.values()), + default=1, + ) -_PAT_COLORS = {"s": "green", "p": "yellow", "u": "red", "uk": "bright_black"} + table = Table( + title=title, + show_header=True, + header_style="bold", + box=None, + padding=(0, 1), + expand=False, + ) + table.add_column("Pattern", width=36, no_wrap=True) + table.add_column("S/P/U/Unk", width=20, no_wrap=True) + table.add_column("", no_wrap=False) -def _render_pattern_matching( - console: Console, - ep_patterns: dict[str, dict[str, dict]], -) -> None: - """Render the PATTERN MATCHING section — per-EP pattern support.""" - if not any(ep_patterns.values()): - return + agg: dict[str, int] = { + "supported": 0, + "partial": 0, + "unsupported": 0, + "unknown": 0, + } - console.print("═" * 80) - console.print("🔍 [bold]PATTERN MATCHING[/bold]") - console.print("═" * 80) + for pattern_id in display_order: + total = all_patterns.get(pattern_id, 0) if all_patterns else sum(data.get(pattern_id, {}).values()) + counts = data.get(pattern_id) - for ep_name, patterns in ep_patterns.items(): - if not patterns: + if not counts: + bar_width = max(1, round(total / max_count * MAX_BAR_WIDTH)) if max_count else 1 + table.add_row( + Text(f" {pattern_id} ({total})", style="dim"), + Text("...", style="dim"), + Text("░" * bar_width, style="dim"), + ) continue - console.print(f" 💻 [bold cyan]{ep_name}[/bold cyan]") + analyzed_for_pattern = sum(counts.get(level, 0) for level in _SUPPORT_LEVEL_KEYS) + for level in agg: + agg[level] += counts.get(level, 0) - for pat_id, pat_info in sorted(patterns.items(), key=lambda x: x[1]["count"], reverse=True): - status = pat_info["status"] - count = pat_info["count"] - icon = _STATUS_ICONS.get(status, "❓") - label = _PATTERN_STATUS_LABELS.get(status, "unknown") - console.print( - f" {icon} [cyan]{pat_id}[/cyan] [dim]({count} instances)[/dim]" - f" — [{_PAT_COLORS.get(status, 'dim')}]{label}[/{_PAT_COLORS.get(status, 'dim')}]" - ) + icon = _worst_level_icon(counts) + pattern_label = Text() + pattern_label.append(f"{icon} ") + pattern_label.append(pattern_id, style="cyan") + if analyzed_for_pattern < total: + pattern_label.append(f" ({analyzed_for_pattern}/{total})", style="dim") + else: + pattern_label.append(f" ({total})", style="dim") - console.print() + bar = _build_stacked_bar(counts, max_count) + remaining = total - analyzed_for_pattern + if remaining > 0: + remaining_width = max(1, round(remaining / max_count * MAX_BAR_WIDTH)) + bar.append("░" * remaining_width, style="dim") + table.add_row(pattern_label, _build_support_text(counts), bar) -def _extract_ep_patterns( - results: list, -) -> dict[str, dict[str, dict]]: - """Extract per-EP subgraph pattern support from analysis results. + table.add_section() + total_patterns = ( + sum(all_patterns.values()) + if all_patterns + else sum(agg.get(level, 0) for level in _SUPPORT_LEVEL_KEYS) + ) + analyzed_count = sum(agg.get(level, 0) for level in _SUPPORT_LEVEL_KEYS) - Args: - results: List of EPSupport objects from AnalysisOutput. + total_label = Text() + total_label.append("TOTAL", style="bold") + if analyzed_count < total_patterns: + total_label.append(f" ({analyzed_count}/{total_patterns})", style="dim") + else: + total_label.append(f" ({total_patterns})", style="dim") + + total_bar = _build_stacked_bar(agg, max(total_patterns, 1)) + total_remaining = total_patterns - analyzed_count + if total_remaining > 0: + total_remaining_width = max( + 1, + round(total_remaining / max(total_patterns, 1) * MAX_BAR_WIDTH), + ) + total_bar.append("░" * total_remaining_width, style="dim") + + table.add_row( + total_label, + _build_support_text(agg), + total_bar, + ) + + return table - Returns: - Dict keyed by EP name, containing dicts of pattern_id to - ``{"count": int, "status": str}`` where status is one of - ``"s"`` (supported), ``"p"`` (partial), ``"u"`` (unsupported), - ``"uk"`` (unknown). - """ - ep_patterns: dict[str, dict[str, dict]] = {} - for ep_support in results: - patterns: dict[str, dict] = {} - for info in ep_support.information: - if info.pattern_id and info.pattern_id.startswith("SUBGRAPH/"): - status = ( - _SUPPORT_LEVEL_TO_SHORT.get(info.status.value, "uk") if info.status else "uk" - ) - patterns[info.pattern_id] = { - "count": len(info.pattern_node_list), - "status": status, - } - ep_patterns[ep_support.ep_type] = patterns - return ep_patterns + +_PATTERN_STATUS_ICONS = { + "supported": "🟢", + "partial": "🟡", + "unsupported": "🔴", + "unknow": "🔵", +} + + +def _pattern_status_view_for_summary( + ep_patterns: dict[str, dict[str, Any]] | None, +) -> dict[str, dict[str, dict[str, Any]]]: + """Normalize pattern payload into {ep: {pattern_id: {count,status}}} view.""" + if not ep_patterns or not isinstance(ep_patterns, dict): + return {} + + summary_view: dict[str, dict[str, dict[str, Any]]] = {} + for ep_name, payload in ep_patterns.items(): + pattern_items = payload.get("patterns", []) if isinstance(payload, dict) else [] + summary_view[ep_name] = { + str(item.get("pattern_id", "")): { + "count": int(item.get("instances", 0)), + "status": str(item.get("status", "unknow")), + } + for item in pattern_items + if str(item.get("pattern_id", "")) + } + + return summary_view def _render_analysis_summary( console: Console, results: list, ep_instance_counts: dict[tuple[str, str], dict[str, dict[str, int]]], - ep_patterns: dict[str, dict[str, dict]] | None = None, + ep_patterns: dict[str, dict[str, Any]] | None = None, *, ep: EPNameOrAlias | Literal["all", "auto"] | None = None, device: str | None = None, no_data_eps: set[tuple[str, str]] | None = None, op_check_skipped: bool = False, + analyze_elapsed_ms: int | None = None, ) -> None: """Render the Analysis Summary section after pattern detection. @@ -377,13 +507,40 @@ def _render_analysis_summary( unknown-op probing). When True, the per-op classification list is suppressed — every op would land in "unknown" with no actionable information. + analyze_elapsed_ms: End-to-end analyze call duration for the current + EP/device run. Rendered as a dim annotation beside the heading. """ from ..analyze.models.support_level import SupportLevel console.print("═" * 80) - console.print("\U0001f4c8 [bold]ANALYSIS SUMMARY[/bold]") + summary_title = "\U0001f4c8 [bold]ANALYSIS SUMMARY[/bold]" + if analyze_elapsed_ms is not None: + if ep is not None and device: + ep_display = _ep_name_device_display_name(str(ep), str(device)) + elif ep is not None: + ep_display = str(ep) + elif results: + first_ep = results[0] + first_ep_name = str(getattr(first_ep, "ep_type", "")) + first_device = str(getattr(first_ep, "device_type", "")).upper() + ep_display = ( + _ep_name_device_display_name(first_ep_name, first_device) + if first_ep_name and first_device + else first_ep_name or "current EP" + ) + else: + ep_display = "current EP" + + elapsed_seconds = max(0.0, analyze_elapsed_ms / 1000.0) + summary_title += ( + f" [dim](Analyze total: {ep_display}, {elapsed_seconds:.2f}s)[/dim]" + ) + + console.print(summary_title) console.print("═" * 80) + pattern_status_view = _pattern_status_view_for_summary(ep_patterns) + if not results: ep_label: str = ep or "all EPs" if device: @@ -414,7 +571,8 @@ def _render_analysis_summary( ep_data = {} has_instance_data = any( sum( - counts.get(level, 0) for level in ("supported", "partial", "unsupported", "unknown") + counts.get(level, 0) + for level in _SUPPORT_LEVEL_KEYS ) > 0 for counts in ep_data.values() @@ -423,14 +581,14 @@ def _render_analysis_summary( # For EPs with no rule data, skip op-level rows — only show patterns. # Always render at least a header so the EP is visible in the summary. if no_data_eps and ep_device_pair in no_data_eps and not has_instance_data: - patterns = (ep_patterns or {}).get(ep_name, {}) + patterns = pattern_status_view.get(ep_name, {}) console.print(f" 🔵 [bold bright_black]{ep_label}[/bold bright_black]:") if patterns: console.print(" [dim]Op check skipped — no rule data[/dim]") for pid, p in sorted(patterns.items(), key=lambda x: x[1]["count"], reverse=True): status = p["status"] - icon_p = _STATUS_ICONS.get(status, "❓") - label = _PATTERN_STATUS_LABELS.get(status, "unknown") + icon_p = _PATTERN_STATUS_ICONS.get(status, "❓") + label = status console.print( f" {icon_p} [dim]{pid}[/dim] ({p['count']} instances, {label})" ) @@ -439,7 +597,12 @@ def _render_analysis_summary( console.print() continue - agg: dict[str, int] = {"supported": 0, "partial": 0, "unsupported": 0, "unknown": 0} + agg: dict[str, int] = { + "supported": 0, + "partial": 0, + "unsupported": 0, + "unknown": 0, + } for counts in ep_data.values(): for level in agg: agg[level] += counts.get(level, 0) @@ -451,12 +614,15 @@ def _render_analysis_summary( ep_style = "bold red" elif agg.get("partial", 0) > 0: ep_style = "bold yellow" - elif agg.get("unknown", 0) > 0 and agg.get("supported", 0) == 0: + elif ( + agg.get("unknown", 0) > 0 + and agg.get("supported", 0) == 0 + ): ep_style = "bold bright_black" else: ep_style = "bold green" - analyzed = _build_analyzed_text(agg) + analyzed = _build_support_text(agg) console.print(f" {icon} [{ep_style}]{ep_label}[/{ep_style}]: ", end="") console.print(analyzed) @@ -468,23 +634,28 @@ def _render_analysis_summary( (SupportLevel.UNKNOWN, "bright_black", "\u2753 Unknown"), ] classification = ep_support.classification + visible_op_names = set(ep_data) if not op_check_skipped: for level, color, heading in _issue_sections: - ops = classification.get(level, []) + ops = [ + op + for op in classification.get(level, []) + if _display_name(op) in visible_op_names + ] if ops: console.print(f" [{color}]{heading}:[/{color}]") for op in sorted(ops): console.print(f" \u2022 [dim]{op}[/dim]") # List non-supported patterns for this EP - patterns = (ep_patterns or {}).get(ep_name, {}) - bad_patterns = {pid: p for pid, p in patterns.items() if p["status"] != "s"} + patterns = pattern_status_view.get(ep_name, {}) + bad_patterns = {pid: p for pid, p in patterns.items() if p["status"] != "supported"} if bad_patterns: console.print(" [dim]Patterns:[/dim]") for pid, p in sorted(bad_patterns.items(), key=lambda x: x[1]["count"], reverse=True): status = p["status"] - icon_p = _STATUS_ICONS.get(status, "\u2753") - label = _PATTERN_STATUS_LABELS.get(status, "unknown") + icon_p = _PATTERN_STATUS_ICONS.get(status, "\u2753") + label = status console.print( f" {icon_p} [dim]{pid}[/dim] ({p['count']} instances, {label})" ) @@ -492,7 +663,12 @@ def _render_analysis_summary( # "Ready to deploy" requires actual op-check data; suppress when skipped. if not op_check_skipped: has_issues = ( - any(classification.get(lvl) for lvl, _, _ in _issue_sections) or bad_patterns + any( + _display_name(op) in visible_op_names + for lvl, _, _ in _issue_sections + for op in classification.get(lvl, []) + ) + or bad_patterns ) if not has_issues: console.print(" [green]Ready to deploy[/green]") @@ -1227,7 +1403,7 @@ def analyze( if not quiet: console.print() console.print("═" * 80) - console.print("📊 [bold]OP CHECK[/bold]") + console.print("📊 [bold]ANALYSIS PROGRESS[/bold]") console.print("═" * 80) console.print(f" 📦 Model: [bold cyan]{model.name}[/bold cyan]") @@ -1265,12 +1441,16 @@ def analyze( current_device = execution_pairs[0][1] all_op_counts: dict[str, int] = {} instance_counts: dict[str, dict[str, int]] = {} + all_pattern_counts: dict[str, int] = {} + pattern_instance_counts: dict[str, dict[str, int]] = {} ep_instance_counts: dict[tuple[str, str], dict[str, dict[str, int]]] = {} live: Live | None = None + pattern_live: Live | None = None unknown_op_progress: Progress | None = None unknown_op_task_id: TaskID | None = None unknown_op_total_nodes = 0 ep_counter = 0 + ep_header_rendered = False _no_data_eps: set[tuple[str, str]] = set() # EP/device pairs with no op rule data analysis_results: list = [] current_run_unknown_op = False @@ -1320,6 +1500,27 @@ def _finalize_unknown_op_progress() -> None: unknown_op_task_id = None unknown_op_total_nodes = 0 + def _finalize_pattern_live(mark_complete: bool = True) -> None: + """Stop active pattern-query Live display, optionally marking complete.""" + nonlocal pattern_live + if pattern_live is None: + return + try: + if mark_complete: + pattern_live.update( + _build_pattern_query_table( + pattern_instance_counts, + ep_device_pair_display_name=_current_ep_device_pair_display_name(), + complete=True, + all_patterns=all_pattern_counts, + ) + ) + except Exception: + logger.debug("Failed to render final pattern table", exc_info=True) + finally: + pattern_live.stop() + pattern_live = None + def _finalize_live(mark_complete: bool = True) -> None: """Stop the active Live display, optionally marking it complete.""" nonlocal live @@ -1345,22 +1546,106 @@ def _finalize_live(mark_complete: bool = True) -> None: live.stop() live = None + def on_pattern_query_start(ep_name: EPName, pattern_counts: dict[str, int]) -> None: + """Called when pattern query stage starts for one EP.""" + nonlocal current_ep_device_pair + nonlocal pattern_instance_counts, all_pattern_counts, ep_counter, pattern_live + nonlocal ep_header_rendered + + # Safety: finalize any stale displays. + _finalize_pattern_live() + _finalize_live() + _finalize_unknown_op_progress() + + current_ep_device_pair = (ep_name, current_device) + all_pattern_counts = { + str(pattern_id): int(total) + for pattern_id, total in pattern_counts.items() + if int(total) > 0 + } + pattern_instance_counts = {} + + ep_counter += 1 + console.print("─" * 80) + console.print( + f"💻 [bold]EP {ep_counter}[/bold]: [bold cyan]{ep_name}[/bold cyan] " + f"on [bold]{current_device}[/bold]" + ) + console.print("─" * 80) + ep_header_rendered = True + + pattern_live = Live( + _build_pattern_query_table( + pattern_instance_counts, + ep_device_pair_display_name=_current_ep_device_pair_display_name(), + all_patterns=all_pattern_counts, + ), + console=console, + refresh_per_second=30, + ) + pattern_live.start() + + def on_pattern_query_result(ep_name: EPName, pattern_id: str, support_status: str) -> None: + """Called when one pattern instance gets a query status.""" + if current_ep_device_pair is None: + return + if ep_name != current_ep_device_pair[0]: + return + + status = str(support_status).strip().lower() + if status == "unknow": + status = "unknown" + if status not in _SUPPORT_LEVEL_KEYS: + status = "unknown" + + counts = pattern_instance_counts.setdefault(str(pattern_id), {}) + counts[status] = counts.get(status, 0) + 1 + + if pattern_live is not None: + pattern_live.update( + _build_pattern_query_table( + pattern_instance_counts, + ep_device_pair_display_name=_current_ep_device_pair_display_name(), + all_patterns=all_pattern_counts, + ) + ) + + def on_pattern_summary_ready(ep_name: EPName, ep_payload: dict[str, Any]) -> None: + """Finalize pattern progress display before OP CHECK starts.""" + _ = ep_name + _finalize_pattern_live() + console.print() + console.print(_build_pattern_coverage_op_line(ep_payload), soft_wrap=True) + def on_ep_start(ep_name: EPName, operator_counts: dict[str, int]) -> None: - """Called when analysis starts for a new EP.""" + """Called when OP CHECK stage starts for a new EP.""" nonlocal current_ep_device_pair - nonlocal instance_counts, all_op_counts, ep_counter, live + nonlocal instance_counts, all_op_counts, live nonlocal unknown_op_progress, unknown_op_task_id, unknown_op_total_nodes nonlocal current_run_unknown_op, current_op_check_skipped + nonlocal ep_counter, ep_header_rendered - # Finalize previous EP's Live display - if current_ep_device_pair is not None: - _finalize_live() - _finalize_unknown_op_progress() - console.print() # blank line between EP tables + _finalize_pattern_live() + _finalize_live() + _finalize_unknown_op_progress() + + if not ep_header_rendered: + ep_counter += 1 + console.print("─" * 80) + console.print( + f"💻 [bold]EP {ep_counter}[/bold]: [bold cyan]{ep_name}[/bold cyan] " + f"on [bold]{current_device}[/bold]" + ) + console.print("─" * 80) + ep_header_rendered = True # Reset for new EP (normalize keys to display names) current_ep_device_pair = (ep_name, current_device) - all_op_counts = {_display_name(k): v for k, v in operator_counts.items()} + all_op_counts = { + _display_name(k): int(v) + for k, v in operator_counts.items() + if int(v) > 0 + } instance_counts = {} has_rule_data = has_rule_data_for_ep(ep_name, current_device) @@ -1373,16 +1658,16 @@ def on_ep_start(ep_name: EPName, operator_counts: dict[str, int]) -> None: _no_data_eps.add((ep_name, current_device)) if current_run_unknown_op: - ep_counter += 1 total_nodes = sum(operator_counts.values()) unknown_op_total_nodes = max(0, total_nodes) - console.print("─" * 80) - console.print( - f"💻 [bold]EP {ep_counter}[/bold]: [bold cyan]{ep_name}[/bold cyan] " - f"on [bold]{current_device}[/bold]" - ) - console.print("─" * 80) + if unknown_op_total_nodes == 0: + console.print( + " [green]All operators are covered by pattern matching; " + "no OP CHECK nodes remain.[/green]" + ) + return + console.print( " [yellow]No rule data detected; probing unknown ops " "one by one...[/yellow]" @@ -1402,15 +1687,7 @@ def on_ep_start(ep_name: EPName, operator_counts: dict[str, int]) -> None: ) return - ep_counter += 1 - - # EP section header - console.print("─" * 80) - console.print( - f"💻 [bold]EP {ep_counter}[/bold]: [bold cyan]{ep_name}[/bold cyan] " - f"on [bold]{current_device}[/bold]" - ) - console.print("─" * 80) + console.print() # Start new Live display — all ops shown as pending live = Live( @@ -1427,6 +1704,10 @@ def on_ep_start(ep_name: EPName, operator_counts: dict[str, int]) -> None: def on_node_result(pattern_runtime: PatternRuntime) -> None: """Callback invoked per-node during analysis.""" + if pattern_runtime.result.reason == "pattern_matched": + # Pattern-matched nodes are excluded from OP CHECK totals and rows. + return + op = _display_name(pattern_runtime.pattern_id) level = pattern_runtime.result.classification.value op_counts = instance_counts.setdefault(op, {}) @@ -1465,6 +1746,7 @@ def on_node_result(pattern_runtime: PatternRuntime) -> None: for target_ep, target_device in execution_pairs: current_device = target_device current_ep_device_pair = None + ep_header_rendered = False run_unknown_op_for_ep = _resolve_run_unknown_op( target_ep, target_device, run_unknown_op, local_pairs @@ -1472,6 +1754,7 @@ def on_node_result(pattern_runtime: PatternRuntime) -> None: current_run_unknown_op = run_unknown_op_for_ep + analyze_start = time.perf_counter() result = analyzer.analyze( model_path=str(model), ep=target_ep, @@ -1483,11 +1766,14 @@ def on_node_result(pattern_runtime: PatternRuntime) -> None: save_node_types=save_node_types, on_node_result=on_node_result, on_ep_start=on_ep_start, + on_pattern_query_start=on_pattern_query_start, + on_pattern_query_result=on_pattern_query_result, + on_pattern_summary_ready=on_pattern_summary_ready, ) + analyze_elapsed_ms = int((time.perf_counter() - analyze_start) * 1000) analysis_results.append(result) - # Extract per-EP pattern support (available now) - ep_patterns = _extract_ep_patterns(result.output.results) + ep_patterns = result.pattern_matching_by_ep # Finalize last EP's Live display _finalize_live() @@ -1495,9 +1781,6 @@ def on_node_result(pattern_runtime: PatternRuntime) -> None: console.print() - # Pattern Matching section (per-EP) - _render_pattern_matching(console, ep_patterns) - # Analysis Summary section _render_analysis_summary( console, @@ -1508,6 +1791,7 @@ def on_node_result(pattern_runtime: PatternRuntime) -> None: device=target_device, no_data_eps=_no_data_eps, op_check_skipped=current_op_check_skipped, + analyze_elapsed_ms=analyze_elapsed_ms, ) # Legend (at the very bottom, only when there are EP results) @@ -1539,6 +1823,7 @@ def on_node_result(pattern_runtime: PatternRuntime) -> None: ) finally: # Safety: stop Live if still running (e.g. on exception) + _finalize_pattern_live(mark_complete=False) _finalize_live(mark_complete=False) _finalize_unknown_op_progress() root_logger.handlers = old_handlers diff --git a/src/winml/modelkit/pattern/base.py b/src/winml/modelkit/pattern/base.py index cac928416..687600f2c 100644 --- a/src/winml/modelkit/pattern/base.py +++ b/src/winml/modelkit/pattern/base.py @@ -931,7 +931,8 @@ def get_onnx_model( output_dtypes[output_idx] ).tensor_proto_type - output_tensor = helper.make_tensor_value_info(output_name, elem_type, None) + # Keep shape present for ONNX checker while leaving dimensions unknown. + output_tensor = helper.make_tensor_value_info(output_name, elem_type, [None]) graph_outputs.append(output_tensor) # Create graph diff --git a/tests/unit/analyze/core/test_pattern_deduplication.py b/tests/unit/analyze/core/test_pattern_deduplication.py index b2d855202..722f0743b 100644 --- a/tests/unit/analyze/core/test_pattern_deduplication.py +++ b/tests/unit/analyze/core/test_pattern_deduplication.py @@ -62,13 +62,9 @@ class TestPatternDeduplication: """Test pattern deduplication logic.""" def test_deduplication_removes_duplicates(self, simple_model_with_tags: ONNXModel, monkeypatch): - """Test that duplicate matches are removed based on node sets.""" + """Test that EP-priority dedup removes duplicate node-key matches in summary path.""" extractor = PatternExtractor(simple_model_with_tags) - # Create two matches with same nodes - div_node = helper.make_node("Div", ["input"], ["div_out"], name="div1") - erf_node = helper.make_node("Erf", ["div_out"], ["output"], name="erf1") - pattern = SubgraphPattern( pattern_id="SUBGRAPH/Test", pattern_name="Test", @@ -76,71 +72,60 @@ def test_deduplication_removes_duplicates(self, simple_model_with_tags: ONNXMode edge_topology=[("div", "erf")], ) - skeleton1 = SkeletonMatchResult( - pattern=pattern, - matched_nodes=[div_node, erf_node], - matched_node_keys=_stable_test_node_keys([div_node, erf_node]), - matcher=None, - ) + def make_match(source: str) -> PatternMatchResult: + div_node = helper.make_node("Div", ["input"], ["div_out"], name=f"div1_{source}") + erf_node = helper.make_node("Erf", ["div_out"], ["output"], name=f"erf1_{source}") + skeleton = SkeletonMatchResult( + pattern=pattern, + matched_nodes=[div_node, erf_node], + matched_node_keys=["stable_dup_node_a", "stable_dup_node_b"], + matcher=None, + ) + return PatternMatchResult( + skeleton_match_result=skeleton, + schema_input_to_value={}, + schema_output_to_value={}, + type_param_to_type={}, + attributes={"source": source}, + ) + + default_match = make_match("default") + ep_match = make_match("qnn") + + def mock_extract(self, *, source, model_signature): + if source == "qnn": + grouped = {"TestPattern": [ep_match]} + else: + grouped = {"TestPattern": [default_match]} + stat = { + "source": source, + "cache_hit": False, + "pattern_class_count": len(grouped), + "match_count": 1, + "elapsed_ms": 0, + } + return grouped, stat - match1 = PatternMatchResult( - skeleton_match_result=skeleton1, - schema_input_to_value={}, - schema_output_to_value={}, - type_param_to_type={}, - attributes={"source": "htp_metadata"}, - ) - - skeleton2 = SkeletonMatchResult( - pattern=pattern, - matched_nodes=[div_node, erf_node], - matched_node_keys=_stable_test_node_keys([div_node, erf_node]), - matcher=None, - ) - - match2 = PatternMatchResult( - skeleton_match_result=skeleton2, - schema_input_to_value={}, - schema_output_to_value={}, - type_param_to_type={}, - attributes={"source": "pattern_matcher"}, + monkeypatch.setattr( + PatternExtractor, + "_extract_skeleton_matches_for_source", + mock_extract, ) - - # Mock the methods to return our test matches - def mock_tag_match(pattern_def): - return [match1] - - def mock_matcher_match(): - return [match2] - monkeypatch.setattr( - extractor, - "_match_subgraph_pattern_from_model_tags", - mock_tag_match, + PatternExtractor, + "_build_merge_prep_metadata", + lambda *args, **kwargs: [], ) monkeypatch.setattr( - extractor, - "extract_subgraph_patterns_with_pattern_matcher", - mock_matcher_match, + PatternExtractor, + "_resolve_sources_for_ep", + lambda *args, **kwargs: ["default", "qnn"], ) - # Mock UnifiedPatternConfig to return test pattern - from unittest.mock import MagicMock, patch - - mock_config = MagicMock() - mock_config.get_htp_patterns.return_value = [pattern] - - with patch( - "winml.modelkit.analyze.core.pattern_extractor.UnifiedPatternConfig", - return_value=mock_config, - ): - # Extract patterns with deduplication - patterns = extractor.extract_subgraph_patterns() - - # Should only have 1 match (duplicate removed) - assert len(patterns) == 1 - # Should keep the first one (from HTP/tag, not PatternMatcher) - assert patterns[0].attributes.get("source") == "htp_metadata" + result = extractor.summary(ep="QNNExecutionProvider", device="NPU") + patterns = result["subgraph_patterns"] + assert len(patterns) == 1 + assert patterns[0] is ep_match def test_different_node_sets_not_deduplicated(self, simple_model_with_tags): """Test that matches with different node sets are kept.""" diff --git a/tests/unit/analyze/core/test_pattern_extractor.py b/tests/unit/analyze/core/test_pattern_extractor.py index b3a1833eb..2b8206cbb 100644 --- a/tests/unit/analyze/core/test_pattern_extractor.py +++ b/tests/unit/analyze/core/test_pattern_extractor.py @@ -6,6 +6,8 @@ from __future__ import annotations +from pathlib import Path +from types import SimpleNamespace from unittest.mock import MagicMock, patch import onnx @@ -14,7 +16,7 @@ from tests.unit.test_helpers import stable_test_node_keys as _stable_test_node_keys from winml.modelkit.analyze import ModelStats, ONNXModel, PatternExtractor -from winml.modelkit.pattern import PatternMatchResult, SkeletonMatchResult, SubgraphPattern +from winml.modelkit.pattern import SubgraphPattern @pytest.fixture @@ -147,196 +149,485 @@ def test_summary_includes_detected_pattern_count( assert result["summary"].detected_pattern_count == {} -class TestPatternExtractorExtractSubgraphPatterns: - """Tests for extract_subgraph_patterns method.""" +class TestPatternExtractorModelSummary: + """Tests for model_summary method.""" - @patch("winml.modelkit.analyze.core.pattern_extractor.UnifiedPatternConfig") - def test_extract_with_no_patterns_returns_empty_list( - self, mock_config_cls: MagicMock, simple_onnx_model: ONNXModel - ) -> None: - """Test extract_subgraph_patterns with no pattern definitions.""" - mock_config = MagicMock() - mock_config.get_htp_patterns.return_value = [] - mock_config_cls.return_value = mock_config + def test_model_summary_returns_metadata(self, simple_onnx_model: ONNXModel) -> None: + """Test model_summary returns ModelStats.""" + extractor = PatternExtractor(simple_onnx_model) + metadata = extractor.model_summary() + + assert isinstance(metadata, ModelStats) + assert metadata.model_path == "test.onnx" + assert metadata.opset_version == 13 + def test_model_summary_with_pattern_count(self, simple_onnx_model: ONNXModel) -> None: + """Test model_summary includes detected_pattern_count.""" extractor = PatternExtractor(simple_onnx_model) - patterns = extractor.extract_subgraph_patterns() + pattern_count_dict = {"SUBGRAPH/GELU_Erf": 5} + metadata = extractor.model_summary(detected_pattern_count=pattern_count_dict) - assert patterns == [] + assert metadata.detected_pattern_count == pattern_count_dict - @patch("winml.modelkit.analyze.core.pattern_extractor.UnifiedPatternConfig") - def test_extract_with_patterns_calls_match( - self, - mock_config_cls: MagicMock, - simple_onnx_model: ONNXModel, - mock_subgraph_pattern: SubgraphPattern, + def test_model_summary_default_pattern_count_is_zero( + self, simple_onnx_model: ONNXModel ) -> None: - """Test extract_subgraph_patterns calls _match_subgraph_pattern_from_model_tags.""" - mock_config = MagicMock() - mock_config.get_htp_patterns.return_value = [mock_subgraph_pattern] - mock_config_cls.return_value = mock_config + """Test model_summary default detected_pattern_count is empty dict.""" + extractor = PatternExtractor(simple_onnx_model) + metadata = extractor.model_summary() + + assert metadata.detected_pattern_count == {} + def test_model_summary_includes_operator_counts(self, simple_onnx_model: ONNXModel) -> None: + """Test model_summary includes operator statistics.""" extractor = PatternExtractor(simple_onnx_model) + metadata = extractor.model_summary() - # Patch _match_subgraph_pattern_from_model_tags to verify it's called - with patch.object( - extractor, "_match_subgraph_pattern_from_model_tags", return_value=[] - ) as mock_match: - patterns = extractor.extract_subgraph_patterns() + assert metadata.total_operators == 2 + assert metadata.unique_operator_types == 2 + assert "Conv" in metadata.operator_counts + assert "Relu" in metadata.operator_counts - mock_match.assert_called_once_with(mock_subgraph_pattern) - assert patterns == [] - @patch("winml.modelkit.analyze.core.pattern_extractor.UnifiedPatternConfig") - def test_extract_returns_matched_patterns( +class TestPatternExtractorAlternativeSelection: + """Tests for merge-prep alternative sorting/filtering helpers.""" + + @staticmethod + def _make_candidate( + *, + pattern_id: str, + pattern_class: str, + is_alternative: bool, + status: str, + compile_ok: bool | None, + run_ok: bool | None, + ) -> dict[str, object]: + return { + "pattern_class": pattern_class, + "pattern_id": pattern_id, + "is_alternative": is_alternative, + "status": status, + "mismatch_error": None, + "compile": compile_ok, + "run": run_ok, + "row_count": 1, + "table_file": "dummy.parquet", + "table_path": "dummy.parquet", + "domain": "ai.onnx", + "opset_version": 17, + "compile_true_rows": int(bool(compile_ok)), + "run_true_rows": int(bool(run_ok)), + "case_indices": None, + "query_condition_count": 0, + "query_condition_keys": [], + "debug_details": None, + } + + def test_select_and_filter_prefers_status_before_priority( self, - mock_config_cls: MagicMock, simple_onnx_model: ONNXModel, - mock_subgraph_pattern: SubgraphPattern, ) -> None: - """Test extract_subgraph_patterns returns matched patterns.""" - mock_config = MagicMock() - mock_config.get_htp_patterns.return_value = [mock_subgraph_pattern] - mock_config_cls.return_value = mock_config - - # Create a mock PatternMatchResult with proper NodeProto objects - from onnx import helper - - # Create mock node protos - conv_node = helper.make_node("Conv", ["input"], ["conv_out"], name="conv1") - relu_node = helper.make_node("Relu", ["conv_out"], ["output"], name="relu1") + """supported status wins even when its priority value is larger.""" + extractor = PatternExtractor(simple_onnx_model) - # Create SkeletonMatchResult - skeleton_result = SkeletonMatchResult( - pattern=mock_subgraph_pattern, - matched_nodes=[conv_node, relu_node], - matched_node_keys=_stable_test_node_keys([conv_node, relu_node]), - matcher=None, + alternatives_meta = [ + {"pattern_to_id": "SUBGRAPH/AltA", "pattern_class": "AltA", "priority": 1}, + {"pattern_to_id": "SUBGRAPH/AltB", "pattern_class": "AltB", "priority": 2}, + ] + candidate_results = [ + self._make_candidate( + pattern_id="SUBGRAPH/Base", + pattern_class="BasePattern", + is_alternative=False, + status="ok", + compile_ok=True, + run_ok=True, + ), + self._make_candidate( + pattern_id="SUBGRAPH/AltA", + pattern_class="AltA", + is_alternative=True, + status="ok", + compile_ok=False, + run_ok=True, + ), + self._make_candidate( + pattern_id="SUBGRAPH/AltB", + pattern_class="AltB", + is_alternative=True, + status="ok", + compile_ok=True, + run_ok=True, + ), + ] + + selected_alternatives, filtered_candidates = extractor._select_and_filter_alternatives( + alternatives_meta=alternatives_meta, + candidate_results=candidate_results, # type: ignore[arg-type] ) - mock_match = PatternMatchResult( - skeleton_match_result=skeleton_result, - schema_input_to_value={}, - schema_output_to_value={}, - type_param_to_type={}, - ) + assert len(selected_alternatives) == 1 + assert selected_alternatives[0]["pattern_to_id"] == "SUBGRAPH/AltB" + + alternative_candidates = [ + candidate for candidate in filtered_candidates if candidate["is_alternative"] + ] + assert len(filtered_candidates) == 2 + assert len(alternative_candidates) == 1 + assert alternative_candidates[0]["pattern_id"] == "SUBGRAPH/AltB" + def test_select_and_filter_uses_priority_as_tiebreaker( + self, + simple_onnx_model: ONNXModel, + ) -> None: + """When statuses tie, lower priority value is selected.""" extractor = PatternExtractor(simple_onnx_model) - with patch.object( - extractor, "_match_subgraph_pattern_from_model_tags", return_value=[mock_match] - ) as _: - patterns = extractor.extract_subgraph_patterns() + alternatives_meta = [ + {"pattern_to_id": "SUBGRAPH/AltA", "pattern_class": "AltA", "priority": 2}, + {"pattern_to_id": "SUBGRAPH/AltB", "pattern_class": "AltB", "priority": 1}, + ] + candidate_results = [ + self._make_candidate( + pattern_id="SUBGRAPH/Base", + pattern_class="BasePattern", + is_alternative=False, + status="ok", + compile_ok=True, + run_ok=True, + ), + self._make_candidate( + pattern_id="SUBGRAPH/AltA", + pattern_class="AltA", + is_alternative=True, + status="ok", + compile_ok=True, + run_ok=True, + ), + self._make_candidate( + pattern_id="SUBGRAPH/AltB", + pattern_class="AltB", + is_alternative=True, + status="ok", + compile_ok=True, + run_ok=True, + ), + ] + + selected_alternatives, filtered_candidates = extractor._select_and_filter_alternatives( + alternatives_meta=alternatives_meta, + candidate_results=candidate_results, # type: ignore[arg-type] + ) + + assert len(selected_alternatives) == 1 + assert selected_alternatives[0]["pattern_to_id"] == "SUBGRAPH/AltB" + assert len(filtered_candidates) == 2 + assert filtered_candidates[1]["pattern_id"] == "SUBGRAPH/AltB" - assert len(patterns) == 1 - assert patterns[0] == mock_match + def test_select_and_filter_drops_selected_unsupported_alternative( + self, + simple_onnx_model: ONNXModel, + ) -> None: + """If the selected best-ranked alternative is unsupported, remove alternatives.""" + extractor = PatternExtractor(simple_onnx_model) + alternatives_meta = [ + {"pattern_to_id": "SUBGRAPH/AltUnsupported", "pattern_class": "AltUnsupported", "priority": 1}, + {"pattern_to_id": "SUBGRAPH/AltUnknown", "pattern_class": "AltUnknown", "priority": 1}, + ] + candidate_results = [ + self._make_candidate( + pattern_id="SUBGRAPH/Base", + pattern_class="BasePattern", + is_alternative=False, + status="ok", + compile_ok=True, + run_ok=True, + ), + self._make_candidate( + pattern_id="SUBGRAPH/AltUnsupported", + pattern_class="AltUnsupported", + is_alternative=True, + status="ok", + compile_ok=False, + run_ok=False, + ), + self._make_candidate( + pattern_id="SUBGRAPH/AltUnknown", + pattern_class="AltUnknown", + is_alternative=True, + status="table_not_found", + compile_ok=None, + run_ok=None, + ), + ] + + selected_alternatives, filtered_candidates = extractor._select_and_filter_alternatives( + alternatives_meta=alternatives_meta, + candidate_results=candidate_results, # type: ignore[arg-type] + ) -class TestPatternExtractorMatchSubgraphPatternFromModelTags: - """Tests for _match_subgraph_pattern_from_model_tags method.""" + assert selected_alternatives == [] + assert len(filtered_candidates) == 1 + assert filtered_candidates[0]["is_alternative"] is False - def test_match_returns_empty_list( + @patch("winml.modelkit.analyze.core.pattern_extractor.UnifiedPatternConfig") + def test_merge_prep_uses_cache_after_first_build( self, + mock_config_cls: MagicMock, simple_onnx_model: ONNXModel, - mock_subgraph_pattern: SubgraphPattern, ) -> None: - """Test _match_subgraph_pattern_from_model_tags returns empty list (mock implementation).""" + """Second call with same cache key should reuse cached merge-prep entries.""" + PatternExtractor._MERGE_PREP_CACHE.clear() extractor = PatternExtractor(simple_onnx_model) - matches = extractor._match_subgraph_pattern_from_model_tags(mock_subgraph_pattern) - # Current implementation is a mock that returns empty list - assert matches == [] + pattern_obj = MagicMock() + pattern_obj.pattern_id = "SUBGRAPH/Base" + pattern_match = MagicMock() + pattern_match.pattern = pattern_obj + pattern_match.match_id = "match_1" + pattern_match.matched_node_keys = ["node_a", "node_b"] + pattern_match.input_infos = {} + pattern_match.attributes = {} -class TestPatternExtractorGetSubgraphPatterns: - """Tests for get_subgraph_patterns method.""" + subgraph_patterns_by_source = { + "default": { + "BasePattern": [pattern_match], + } + } - @patch("winml.modelkit.analyze.core.pattern_extractor.UnifiedPatternConfig") - def test_get_patterns_calls_unified_config( - self, mock_config_cls: MagicMock, simple_onnx_model: ONNXModel - ) -> None: - """Test get_subgraph_patterns calls UnifiedPatternConfig.get_htp_patterns.""" mock_config = MagicMock() - mock_config.get_htp_patterns.return_value = [] + mock_config.get_alternatives.return_value = [ + SimpleNamespace( + pattern_to_id="SUBGRAPH/Alt", + pattern_class="AltPattern", + priority=1, + enabled=True, + module=None, + ) + ] mock_config_cls.return_value = mock_config - extractor = PatternExtractor(simple_onnx_model) - patterns = extractor.get_subgraph_patterns() - - mock_config.get_htp_patterns.assert_called_once() - assert patterns == [] + with ( + patch.object( + PatternExtractor, + "_is_valid_parquet_lookup_target", + return_value=True, + ), + patch.object( + PatternExtractor, + "_probe_candidate_pattern_mismatch", + return_value=(False, None), + ), + patch.object( + PatternExtractor, + "_domain_and_target_opset_for_pattern", + return_value=("ai.onnx", 13), + ), + patch.object( + PatternExtractor, + "_resolve_pattern_rule_table", + return_value=(Path("dummy.parquet"), "ai.onnx", 13), + ), + patch.object( + PatternExtractor, + "_query_pattern_rule_compile_run_for_match", + return_value=("ok", True, True, 1, 1, 1, None, 0, [], None), + ) as mock_query, + ): + first = extractor._build_merge_prep_metadata( + subgraph_patterns_by_source=subgraph_patterns_by_source, + model_signature="sig_1", + ep="QNNExecutionProvider", + device="NPU", + for_debug=True, + ) + assert mock_query.call_count == 2 + + second = extractor._build_merge_prep_metadata( + subgraph_patterns_by_source=subgraph_patterns_by_source, + model_signature="sig_1", + ep="QNNExecutionProvider", + device="NPU", + for_debug=True, + ) + assert mock_query.call_count == 2 + + assert first == second + assert first is not second + PatternExtractor._MERGE_PREP_CACHE.clear() @patch("winml.modelkit.analyze.core.pattern_extractor.UnifiedPatternConfig") - def test_get_patterns_returns_loaded_patterns( + def test_merge_prep_stops_after_first_supported_alternative_by_priority( self, mock_config_cls: MagicMock, simple_onnx_model: ONNXModel, - mock_subgraph_pattern: SubgraphPattern, ) -> None: - """Test get_subgraph_patterns returns patterns from UnifiedPatternConfig.""" - mock_config = MagicMock() - mock_config.get_htp_patterns.return_value = [mock_subgraph_pattern] - mock_config_cls.return_value = mock_config - + """Alternative probing stops once the first priority-ordered supported option is found.""" + PatternExtractor._MERGE_PREP_CACHE.clear() extractor = PatternExtractor(simple_onnx_model) - patterns = extractor.get_subgraph_patterns() - assert len(patterns) == 1 - assert patterns[0] == mock_subgraph_pattern + pattern_obj = MagicMock() + pattern_obj.pattern_id = "SUBGRAPH/Base" + + pattern_match = MagicMock() + pattern_match.pattern = pattern_obj + pattern_match.match_id = "match_short_circuit" + pattern_match.matched_node_keys = ["node_short_a", "node_short_b"] + pattern_match.input_infos = {} + pattern_match.attributes = {} + + subgraph_patterns_by_source = { + "default": { + "BasePattern": [pattern_match], + } + } - @patch("winml.modelkit.analyze.core.pattern_extractor.UnifiedPatternConfig") - def test_get_patterns_returns_empty_list_when_no_rules( - self, mock_config_cls: MagicMock, simple_onnx_model: ONNXModel - ) -> None: - """Test get_subgraph_patterns returns empty list when no patterns found.""" mock_config = MagicMock() - mock_config.get_htp_patterns.return_value = [] + mock_config.get_alternatives.return_value = [ + SimpleNamespace( + pattern_to_id="SUBGRAPH/AltLowPriority", + pattern_class="AltLowPriority", + priority=2, + enabled=True, + module=None, + ), + SimpleNamespace( + pattern_to_id="SUBGRAPH/AltHighPriority", + pattern_class="AltHighPriority", + priority=1, + enabled=True, + module=None, + ), + ] mock_config_cls.return_value = mock_config + def query_side_effect(**kwargs: object) -> tuple[str, bool | None, bool | None, int, int, int, list[object] | None, int, list[str], dict[str, object] | None]: + candidate_name = str(kwargs["candidate_pattern_name"]) + if candidate_name == "AltHighPriority": + return ("ok", True, True, 1, 1, 1, None, 0, [], None) + return ("ok", False, False, 1, 0, 0, None, 0, [], None) + + with ( + patch.object( + PatternExtractor, + "_is_valid_parquet_lookup_target", + return_value=True, + ), + patch.object( + PatternExtractor, + "_probe_candidate_pattern_mismatch", + return_value=(False, None), + ), + patch.object( + PatternExtractor, + "_domain_and_target_opset_for_pattern", + return_value=("ai.onnx", 13), + ), + patch.object( + PatternExtractor, + "_resolve_pattern_rule_table", + return_value=(Path("dummy.parquet"), "ai.onnx", 13), + ), + patch.object( + PatternExtractor, + "_query_pattern_rule_compile_run_for_match", + side_effect=query_side_effect, + ) as mock_query, + ): + entries = extractor._build_merge_prep_metadata( + subgraph_patterns_by_source=subgraph_patterns_by_source, + model_signature="sig_short_circuit", + ep="QNNExecutionProvider", + device="NPU", + for_debug=True, + ) + + queried_candidates = [ + str(call.kwargs["candidate_pattern_name"]) for call in mock_query.call_args_list + ] + assert queried_candidates == ["MagicMock", "AltHighPriority"] + assert entries[0]["alternatives"][0]["pattern_to_id"] == "SUBGRAPH/AltHighPriority" + PatternExtractor._MERGE_PREP_CACHE.clear() + + +class TestPatternExtractorEPDedup: + """Tests for EP-priority dedup and EP-scoped dedup cache.""" + + @staticmethod + def _make_match(node_keys: list[str]) -> MagicMock: + match = MagicMock() + match.matched_node_keys = node_keys + return match + + def test_ep_priority_dedup_prefers_ep_source_over_default( + self, + simple_onnx_model: ONNXModel, + ) -> None: + """EP source should be traversed first and win on overlapping node keys.""" + PatternExtractor._DEDUPED_MATCH_CACHE.clear() extractor = PatternExtractor(simple_onnx_model) - patterns = extractor.get_subgraph_patterns() - - assert patterns == [] + ep_match = self._make_match(["shared_node", "ep_only_node"]) + default_overlap = self._make_match(["shared_node", "default_only_node"]) + default_unique = self._make_match(["default_unique_node"]) -class TestPatternExtractorModelSummary: - """Tests for model_summary method.""" - - def test_model_summary_returns_metadata(self, simple_onnx_model: ONNXModel) -> None: - """Test model_summary returns ModelStats.""" - extractor = PatternExtractor(simple_onnx_model) - metadata = extractor.model_summary() + grouped = { + "default": {"DemoPattern": [default_overlap, default_unique]}, + "qnn": {"DemoPattern": [ep_match]}, + } - assert isinstance(metadata, ModelStats) - assert metadata.model_path == "test.onnx" - assert metadata.opset_version == 13 - - def test_model_summary_with_pattern_count(self, simple_onnx_model: ONNXModel) -> None: - """Test model_summary includes detected_pattern_count.""" - extractor = PatternExtractor(simple_onnx_model) - pattern_count_dict = {"SUBGRAPH/GELU_Erf": 5} - metadata = extractor.model_summary(detected_pattern_count=pattern_count_dict) + deduped_grouped, deduped_flat = extractor._dedup_grouped_matches_for_ep( + subgraph_patterns_by_source=grouped, # type: ignore[arg-type] + sources=["default", "qnn"], + model_signature="sig_ep_priority", + ep="QNNExecutionProvider", + ) - assert metadata.detected_pattern_count == pattern_count_dict + assert deduped_flat == [ep_match, default_unique] + assert deduped_grouped["qnn"]["DemoPattern"] == [ep_match] + assert deduped_grouped["default"]["DemoPattern"] == [default_unique] + PatternExtractor._DEDUPED_MATCH_CACHE.clear() - def test_model_summary_default_pattern_count_is_zero( - self, simple_onnx_model: ONNXModel + def test_ep_dedup_cache_reused_for_same_ep( + self, + simple_onnx_model: ONNXModel, ) -> None: - """Test model_summary default detected_pattern_count is empty dict.""" + """Dedup cache key should ignore device and reuse by model+EP.""" + PatternExtractor._DEDUPED_MATCH_CACHE.clear() extractor = PatternExtractor(simple_onnx_model) - metadata = extractor.model_summary() - assert metadata.detected_pattern_count == {} + first_match = self._make_match(["node_a"]) + grouped_first = { + "default": {"DemoPattern": [first_match]}, + "qnn": {"DemoPattern": []}, + } + + first_grouped, first_flat = extractor._dedup_grouped_matches_for_ep( + subgraph_patterns_by_source=grouped_first, # type: ignore[arg-type] + sources=["default", "qnn"], + model_signature="sig_ep_cache", + ep="QNNExecutionProvider", + ) - def test_model_summary_includes_operator_counts(self, simple_onnx_model: ONNXModel) -> None: - """Test model_summary includes operator statistics.""" - extractor = PatternExtractor(simple_onnx_model) - metadata = extractor.model_summary() + second_match = self._make_match(["node_b"]) + grouped_second = { + "default": {"DemoPattern": [second_match]}, + "qnn": {"DemoPattern": []}, + } + + second_grouped, second_flat = extractor._dedup_grouped_matches_for_ep( + subgraph_patterns_by_source=grouped_second, # type: ignore[arg-type] + sources=["default", "qnn"], + model_signature="sig_ep_cache", + ep="QNNExecutionProvider", + ) - assert metadata.total_operators == 2 - assert metadata.unique_operator_types == 2 - assert "Conv" in metadata.operator_counts - assert "Relu" in metadata.operator_counts + assert first_flat == [first_match] + assert second_flat == [first_match] + assert first_grouped == second_grouped + PatternExtractor._DEDUPED_MATCH_CACHE.clear() class TestPatternExtractorIntegration: @@ -396,12 +687,6 @@ def test_workflow_with_multiple_patterns( extractor = PatternExtractor(simple_onnx_model) - # Extract patterns - patterns = extractor.extract_subgraph_patterns() - - # Should process both patterns (even if no matches found) - assert isinstance(patterns, list) - - # Verify both patterns were loaded - loaded_patterns = extractor.get_subgraph_patterns() - assert len(loaded_patterns) == 2 + # Summary should run end-to-end with multiple pattern definitions + result = extractor.summary() + assert isinstance(result["subgraph_patterns"], list) diff --git a/tests/unit/analyze/core/test_runtime_checker.py b/tests/unit/analyze/core/test_runtime_checker.py index 6a8a11de0..c364a9621 100644 --- a/tests/unit/analyze/core/test_runtime_checker.py +++ b/tests/unit/analyze/core/test_runtime_checker.py @@ -7,8 +7,7 @@ Tests verify: - Correct return types for summary() method -- Correct type annotations for alternatives -- Type safety with PatternRuntime and PatternAlternative +- Type safety with PatternRuntime - Cache reuse for RuntimeCheckerQuery """ @@ -19,21 +18,10 @@ import onnx import pytest -from tests.unit.test_helpers import stable_test_node_keys as _stable_test_node_keys -from winml.modelkit.analyze import ONNXModel, RuntimeChecker, RuntimeTestResult +from winml.modelkit.analyze import ONNXModel, RuntimeChecker from winml.modelkit.analyze.core import runtime_checker_query as runtime_checker_query_module from winml.modelkit.analyze.core.runtime_checker_query import RuntimeCheckerQuery -from winml.modelkit.analyze.models.runtime_checks import ( # Testing internal implementation - AlternativeType, - PatternAlternative, - PatternRuntime, -) -from winml.modelkit.pattern import ( - OperatorPattern, - PatternMatchResult, - PatternType, - SkeletonMatchResult, -) +from winml.modelkit.analyze.models.runtime_checks import PatternRuntime TensorProto = onnx.TensorProto @@ -59,49 +47,15 @@ def simple_onnx_model() -> ONNXModel: return ONNXModel.from_onnx_model(model_def, "test.onnx") -@pytest.fixture -def sample_pattern_match() -> PatternMatchResult: - """Create a sample PatternMatchResult for testing.""" - pattern = OperatorPattern( - pattern_id="OP/ai.onnx/Conv", - pattern_type=PatternType.OPERATOR, - namespace="ai.onnx", - op_type="Conv", - description="Conv operator", - ) - - # Create mock node proto matching the model's inputs - node_proto = helper.make_node("Conv", ["input1"], ["conv_output"], name="conv_node") - - # Create SkeletonMatchResult - skeleton_result = SkeletonMatchResult( - pattern=pattern, - matched_nodes=[node_proto], - matched_node_keys=_stable_test_node_keys([node_proto]), - matcher=None, - ) - - return PatternMatchResult( - skeleton_match_result=skeleton_result, - schema_input_to_value={}, - schema_output_to_value={}, - type_param_to_type={}, - ) - - class TestRuntimeCheckerTypeHints: """Test RuntimeChecker return type correctness.""" - def test_summary_returns_correct_type( - self, simple_onnx_model: ONNXModel, sample_pattern_match: PatternMatchResult - ): + def test_summary_returns_correct_type(self, simple_onnx_model: ONNXModel): """Test that summary() returns dict[str, list[PatternRuntime]].""" - # Initialize with both model and patterns to populate summary checker = RuntimeChecker( ep="QNNExecutionProvider", device="NPU", model=simple_onnx_model, - patterns=[sample_pattern_match], ) result = checker.summary() @@ -115,32 +69,25 @@ def test_summary_returns_correct_type( assert isinstance(value, list) assert all(isinstance(item, PatternRuntime) for item in value) - # Verify expected keys (both should be present since we have model + patterns) - assert "op_runtime_check_result" in result - assert "subgraph_runtime_check_result" in result + assert set(result) == {"op_runtime_check_result"} def test_summary_with_model_only(self, simple_onnx_model: ONNXModel): """Test summary() when initialized with model only.""" - # When initialized with only model, summary() needs patterns parameter checker = RuntimeChecker( ep="QNNExecutionProvider", device="NPU", model=simple_onnx_model, ) - # Pass empty patterns to avoid ValueError - result = checker.summary(patterns=[]) + result = checker.summary() - # Should have both keys, but subgraph will be empty assert isinstance(result, dict) - assert "op_runtime_check_result" in result - assert "subgraph_runtime_check_result" in result + assert set(result) == {"op_runtime_check_result"} # Verify types op_results = result["op_runtime_check_result"] assert isinstance(op_results, list) assert all(isinstance(item, PatternRuntime) for item in op_results) - assert len(result["subgraph_runtime_check_result"]) == 0 def test_op_support_returns_list_of_pattern_runtime(self, simple_onnx_model: ONNXModel): """Test that op_support() returns list[PatternRuntime].""" @@ -159,80 +106,17 @@ def test_op_support_returns_list_of_pattern_runtime(self, simple_onnx_model: ONN # Should have one operator (Add node) assert len(result) > 0 - def test_subgraph_support_returns_list_of_pattern_runtime( - self, sample_pattern_match: PatternMatchResult, simple_onnx_model: ONNXModel - ): - """Test that subgraph_support() returns list[PatternRuntime].""" - # Need model for _lookup_pattern_support - checker = RuntimeChecker( - ep="QNNExecutionProvider", - device="NPU", - model=simple_onnx_model, - patterns=[sample_pattern_match], - ) - - result = checker.subgraph_support() - - # Verify return type - assert isinstance(result, list) - assert all(isinstance(item, PatternRuntime) for item in result) - assert len(result) == 1 - - def test_query_pattern_support_returns_pattern_runtime( - self, sample_pattern_match: PatternMatchResult, simple_onnx_model: ONNXModel - ): - """Test that query_pattern_support() returns PatternRuntime.""" - checker = RuntimeChecker( - ep="QNNExecutionProvider", - device="NPU", - model=simple_onnx_model, - ) - - result = checker.query_pattern_support(sample_pattern_match) - - # Verify return type - assert isinstance(result, PatternRuntime) - assert result.pattern_id == "OP/ai.onnx/Conv" - assert isinstance(result.result, RuntimeTestResult) - assert isinstance(result.alternatives, list) - - def test_alternatives_is_list_of_pattern_alternative( - self, sample_pattern_match: PatternMatchResult, simple_onnx_model: ONNXModel - ): - """Test that PatternRuntime.alternatives is list[PatternAlternative].""" - checker = RuntimeChecker( - ep="QNNExecutionProvider", - device="NPU", - model=simple_onnx_model, - ) - - result = checker.query_pattern_support(sample_pattern_match) - - # Verify alternatives type - assert isinstance(result.alternatives, list) - - # Currently alternatives is empty (not implemented) - # But when implemented, should contain PatternAlternative objects - for alt in result.alternatives: - assert isinstance(alt, PatternAlternative) - assert hasattr(alt, "pattern_id") - assert hasattr(alt, "result") - assert hasattr(alt, "alternative_type") - class TestRuntimeCheckerValidation: """Test RuntimeChecker initialization validation.""" - def test_requires_either_model_or_patterns(self): - """Test that RuntimeChecker requires at least one of model or patterns.""" - with pytest.raises( - ValueError, match="At least one of 'model' or 'patterns' must be provided" - ): + def test_requires_model(self): + """Test that RuntimeChecker requires a model.""" + with pytest.raises(ValueError, match="'model' is required"): RuntimeChecker( ep="QNNExecutionProvider", device="NPU", model=None, - patterns=None, ) def test_requires_non_empty_device(self, simple_onnx_model: ONNXModel): @@ -244,28 +128,6 @@ def test_requires_non_empty_device(self, simple_onnx_model: ONNXModel): model=simple_onnx_model, ) - def test_op_support_requires_model(self, sample_pattern_match: PatternMatchResult): - """Test that op_support() requires model to be provided.""" - checker = RuntimeChecker( - ep="QNNExecutionProvider", - device="NPU", - patterns=[sample_pattern_match], - ) - - with pytest.raises(ValueError, match="op_support\\(\\) requires ONNXModel"): - checker.op_support() - - def test_subgraph_support_requires_patterns(self, simple_onnx_model: ONNXModel): - """Test that subgraph_support() requires patterns when not initialized with them.""" - checker = RuntimeChecker( - ep="QNNExecutionProvider", - device="NPU", - model=simple_onnx_model, - ) - - with pytest.raises(ValueError, match="patterns parameter is required"): - checker.subgraph_support(patterns=None) - class TestRuntimeCheckerIntegration: """Integration tests for RuntimeChecker.""" @@ -284,7 +146,7 @@ def test_full_workflow_with_model(self, simple_onnx_model: ONNXModel): assert all(isinstance(r, PatternRuntime) for r in op_results) # Get summary with empty patterns - summary = checker.summary(patterns=[]) + summary = checker.summary() assert isinstance(summary, dict) assert "op_runtime_check_result" in summary assert len(summary["op_runtime_check_result"]) == len(op_results) @@ -378,102 +240,6 @@ def check_run(self, model_bytes, input_feed): assert {vi.name for vi in single_node_model.graph.input} == {"weight", "input"} assert {init.name for init in single_node_model.graph.initializer} == set() - def test_full_workflow_with_patterns( - self, sample_pattern_match: PatternMatchResult, simple_onnx_model: ONNXModel - ): - """Test complete workflow: initialize with patterns, check subgraph support.""" - # Need model for pattern lookup - checker = RuntimeChecker( - ep="QNNExecutionProvider", - device="NPU", - model=simple_onnx_model, - patterns=[sample_pattern_match], - ) - - # Get subgraph support - subgraph_results = checker.subgraph_support() - assert len(subgraph_results) == 1 - assert all(isinstance(r, PatternRuntime) for r in subgraph_results) - - # Get summary - summary = checker.summary() - assert isinstance(summary, dict) - assert "subgraph_runtime_check_result" in summary - assert len(summary["subgraph_runtime_check_result"]) == 1 - - def test_op_merged_from_subgraph_has_empty_alternatives( - self, simple_onnx_model: ONNXModel, monkeypatch: pytest.MonkeyPatch - ): - """Ops merged from a subgraph pattern must have alternatives=[], not the subgraph's. - - When a node is covered by a matched subgraph pattern, summary() replaces the - op-level result with the subgraph-level result. The subgraph may carry - alternatives (e.g. SingleGeluPattern → GeluPattern), but those belong to the - subgraph entry — not to the individual op row. Leaking them onto the op - would misrepresent what alternatives are available for that specific node. - """ - checker = RuntimeChecker( - ep="QNNExecutionProvider", - device="NPU", - model=simple_onnx_model, - ) - - shared_node = helper.make_node("Add", ["a", "b"], ["c"], name="shared_node") - - def _make_pm(node): - pattern = OperatorPattern( - pattern_id=f"OP/ai.onnx/{node.op_type}", - pattern_type=PatternType.OPERATOR, - namespace="ai.onnx", - op_type=node.op_type, - description="", - ) - skeleton = SkeletonMatchResult( - pattern=pattern, - matched_nodes=[node], - matched_node_keys=_stable_test_node_keys([node]), - matcher=None, - ) - return PatternMatchResult( - skeleton_match_result=skeleton, - schema_input_to_value={}, - schema_output_to_value={}, - type_param_to_type={}, - ) - - supported_result = RuntimeTestResult(compile=True, run=True) - subgraph_alternative = PatternAlternative( - pattern_id="SUBGRAPH/SingleGeluPattern", - result=supported_result, - alternative_type=AlternativeType.EQUIVALENT, - ) - - op_pr = PatternRuntime( - pattern_id="OP/ai.onnx/Add", - result=supported_result, - alternatives=[], - pattern_match=_make_pm(shared_node), - ) - subgraph_pr = PatternRuntime( - pattern_id="SUBGRAPH/GeluPattern", - result=supported_result, - alternatives=[subgraph_alternative], # subgraph has a non-empty alternative - pattern_match=_make_pm(shared_node), - ) - - monkeypatch.setattr(checker, "op_support", lambda **kw: [op_pr]) - monkeypatch.setattr(checker, "subgraph_support", lambda *a, **kw: [subgraph_pr]) - - result = checker.summary(patterns=[]) - merged_ops = result["op_runtime_check_result"] - - assert len(merged_ops) == 1 - merged = merged_ops[0] - # Result must be taken from the subgraph - assert merged.result is subgraph_pr.result - # alternatives must be empty — subgraph alternatives must NOT leak onto the op - assert merged.alternatives == [] - class TestRuntimeCheckerQueryCache: """Test RuntimeCheckerQuery caching functionality.""" @@ -500,28 +266,6 @@ def test_query_cache_reuse(self, simple_onnx_model: ONNXModel): # Results should be consistent assert len(first_result) == len(second_result) - def test_query_cache_across_methods( - self, simple_onnx_model: ONNXModel, sample_pattern_match: PatternMatchResult - ): - """Test that query cache is shared across op_support and pattern lookup.""" - checker = RuntimeChecker( - ep="QNNExecutionProvider", - device="NPU", - model=simple_onnx_model, - patterns=[sample_pattern_match], - ) - - # Call op_support first - checker.op_support() - query_after_op_support = checker._query - - # Call query_pattern_support - checker.query_pattern_support(sample_pattern_match) - query_after_pattern_support = checker._query - - # Should be the same cached query - assert query_after_pattern_support is query_after_op_support - def test_query_cache_performance(self, simple_onnx_model: ONNXModel): """Test that cache improves performance on repeated calls.""" checker = RuntimeChecker( @@ -534,29 +278,18 @@ def test_query_cache_performance(self, simple_onnx_model: ONNXModel): start_time = time.time() checker.op_support() _first_call_time = time.time() - start_time + first_query = checker._query # Second call - warm (uses cache) start_time = time.time() checker.op_support() _second_call_time = time.time() - start_time + second_query = checker._query # Second call should be faster or at least not significantly slower # We're primarily checking that it doesn't recreate the query # which would add initialization overhead - assert checker._query is not None # Not asserting timing directly as it can be flaky, # but verifying cache exists proves the optimization - - def test_get_query_without_model_raises_error(self, sample_pattern_match: PatternMatchResult): - """Test that _get_query raises error when model is not available.""" - checker = RuntimeChecker( - ep="QNNExecutionProvider", - device="NPU", - patterns=[sample_pattern_match], - ) - - # _get_query should raise ValueError - with pytest.raises( - ValueError, match="Cannot create RuntimeCheckerQuery without ONNX model" - ): - checker._get_query() + assert first_query is not None + assert second_query is first_query diff --git a/tests/unit/analyze/core/test_runtime_checker_query_parquet.py b/tests/unit/analyze/core/test_runtime_checker_query_parquet.py index f8331ac17..80b661d7e 100644 --- a/tests/unit/analyze/core/test_runtime_checker_query_parquet.py +++ b/tests/unit/analyze/core/test_runtime_checker_query_parquet.py @@ -119,6 +119,47 @@ def clear_debug_rules_env(monkeypatch: pytest.MonkeyPatch): class TestRuntimeCheckerQueryParquet: """Validate parquet runtime rule lookup.""" + def test_pattern_matched_node_skips_parquet_lookup( + self, + tmp_path: Path, + monkeypatch: pytest.MonkeyPatch, + ): + """Nodes covered by pattern hashset should bypass parquet table checks.""" + monkeypatch.setenv("WINMLCLI_RULES_DIR", str(tmp_path)) + + model = _build_add_model() + node = model.graph.node[0] + + def _unexpected_get_conditions(*args, **kwargs): + del args, kwargs + raise AssertionError("get_query_conditions_for_node should not be called") + + monkeypatch.setattr( + runtime_checker_query_module, + "get_query_conditions_for_node", + _unexpected_get_conditions, + ) + + query_parquet = RuntimeCheckerQuery( + model, + "QNNExecutionProvider", + "NPU", + pattern_matched_node_status_by_key={"add_node": "unsupported"}, + ) + query_parquet.node_checkers = [] + + result = query_parquet.run_for_node(node, for_debug=True, run_unknown_op=False) + + assert result.pattern_id == "OP/ai.onnx/Add" + assert result.result.no_data is False + assert result.result.compile is False + assert result.result.run is False + assert result.result.reason == "pattern_matched" + debug_details = result.result.debug_details + assert isinstance(debug_details, dict) + assert debug_details.get("type") == "pattern_matched" + assert debug_details.get("status") == "unsupported" + def test_parquet_lookup_returns_expected_result( self, tmp_path: Path, diff --git a/tests/unit/analyze/test_analyzer.py b/tests/unit/analyze/test_analyzer.py index 43dc86775..5c81dc2a0 100644 --- a/tests/unit/analyze/test_analyzer.py +++ b/tests/unit/analyze/test_analyzer.py @@ -764,18 +764,16 @@ def test_build_runtime_debug_details_summary_groups_and_records_unknown(self) -> }, ), ), - ], - "subgraph_runtime_check_result": [ PatternRuntime( - pattern_id="SUBGRAPH/TestPattern", + pattern_id="OP/ai.onnx/Unsupported", result=RuntimeTestResult( compile=False, run=False, debug_details={ - "node_stable_key": "node_subgraph", + "node_stable_key": "node_unsupported", "case_indices": ["case_5"], - "table_path": "rules/subgraph.parquet", - "table_file": "subgraph.parquet", + "table_path": "rules/unsupported.parquet", + "table_file": "unsupported.parquet", }, ), ) @@ -797,9 +795,9 @@ def test_build_runtime_debug_details_summary_groups_and_records_unknown(self) -> assert summary["partial"]["node_resize"].table_path == "rules/resize.parquet" assert summary["partial"]["node_resize"].table_file == "resize.parquet" - assert summary["unsupported"]["node_subgraph"].case_indices == ["case_5"] - assert summary["unsupported"]["node_subgraph"].table_path == "rules/subgraph.parquet" - assert summary["unsupported"]["node_subgraph"].table_file == "subgraph.parquet" + assert summary["unsupported"]["node_unsupported"].case_indices == ["case_5"] + assert summary["unsupported"]["node_unsupported"].table_path == "rules/unsupported.parquet" + assert summary["unsupported"]["node_unsupported"].table_file == "unsupported.parquet" # Unknown nodes are recorded as a plain list of node keys (no case data). assert summary["unknown"] == ["node_unknown"] @@ -835,7 +833,6 @@ def test_build_runtime_debug_details_summary_merges_same_node(self) -> None: ), ), ], - "subgraph_runtime_check_result": [], } summary = _build_runtime_debug_details_summary(runtime_summary) @@ -921,7 +918,6 @@ def test_analyze_from_proto_single_ep( mock_checker = MagicMock() mock_checker.summary.return_value = { "op_runtime_check_result": [], - "subgraph_runtime_check_result": [], } mock_runtime_checker_cls.return_value = mock_checker @@ -1070,7 +1066,6 @@ def test_analyze_from_proto_includes_runtime_debug_summary_when_debug_enabled( ), ), ], - "subgraph_runtime_check_result": [], } mock_runtime_checker_cls.return_value = mock_checker @@ -1134,7 +1129,6 @@ def test_analyze_from_proto_multi_ep( mock_checker = MagicMock() mock_checker.summary.return_value = { "op_runtime_check_result": [], - "subgraph_runtime_check_result": [], } mock_runtime_checker_cls.return_value = mock_checker @@ -1201,7 +1195,6 @@ def test_analyze_from_proto_default_driver( mock_checker = MagicMock() mock_checker.summary.return_value = { "op_runtime_check_result": [], - "subgraph_runtime_check_result": [], } mock_runtime_checker_cls.return_value = mock_checker @@ -1265,7 +1258,6 @@ def test_analyze_from_proto_with_information( mock_checker.summary.return_value = { "op_runtime_check_result": [mock_pattern_runtime], # Non-empty - "subgraph_runtime_check_result": [], } mock_runtime_checker_cls.return_value = mock_checker @@ -1335,7 +1327,6 @@ def test_analyze_from_proto_always_runs_ep( mock_runtime_checker = MagicMock() mock_runtime_checker.summary.return_value = { "op_runtime_check_result": [], - "subgraph_runtime_check_result": [], } mock_runtime_checker_cls.return_value = mock_runtime_checker diff --git a/tests/unit/analyze/test_static_analyzer_cli.py b/tests/unit/analyze/test_static_analyzer_cli.py index b7a078dd3..ac3d8e92e 100644 --- a/tests/unit/analyze/test_static_analyzer_cli.py +++ b/tests/unit/analyze/test_static_analyzer_cli.py @@ -1971,6 +1971,32 @@ def invoke_callbacks(**kwargs): class TestAnalyzeSummaryRendering: """Summary rendering behavior for no-rule-data fallback cases.""" + def test_summary_heading_includes_per_ep_analyze_elapsed(self) -> None: + """Heading should show elapsed analyze time annotation for EP/device.""" + from winml.modelkit.commands.analyze import _render_analysis_summary + + console = Console(record=True, force_terminal=False, width=120) + + ep_support = Mock() + ep_support.ep_type = "DmlExecutionProvider" + ep_support.device_type = "GPU" + ep_support.classification = {} + ep_support.information = [] + + _render_analysis_summary( + console, + [ep_support], + ep_instance_counts={("DmlExecutionProvider", "GPU"): {"Conv": {"supported": 1}}}, + ep_patterns={}, + ep="DmlExecutionProvider", + device="GPU", + analyze_elapsed_ms=1234, + ) + + output = console.export_text() + assert "ANALYSIS SUMMARY" in output + assert "Analyze total: DmlExecutionProvider (GPU), 1.23s" in output + def test_no_rule_data_with_instance_counts_renders_op_summary(self) -> None: """When unknown-op probing produced counts, summary should not show skip message.""" from winml.modelkit.commands.analyze import _render_analysis_summary