Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
62 changes: 61 additions & 1 deletion src/winml/modelkit/analyze/optim_output.py
Original file line number Diff line number Diff line change
Expand Up @@ -24,6 +24,7 @@
from dataclasses import dataclass, field
from typing import TYPE_CHECKING, cast

from ..optim.analysis import NodeRef
from .models.onnx_model import ONNXModel
from .models.support_level import SupportLevel

Expand All @@ -33,7 +34,7 @@

from onnx import ModelProto, NodeProto

from ..optim import CapabilityFinding, NodeRef
from ..optim import CapabilityFinding
from ..utils.constants import EPName


Expand Down Expand Up @@ -68,6 +69,18 @@ class ProducedOperatorSupport:
support: SupportLevel
reason: str | None = None

def to_dict(self) -> dict[str, str]:
"""Return the stable JSON representation of this produced operator."""
data = {
"op_type": self.op_type,
"label": self.label,
"change": self.change,
"support": self.support.value,
}
if self.reason is not None:
data["reason"] = self.reason
return data


@dataclass
class OptimizationOutputSupport:
Expand All @@ -88,6 +101,12 @@ class OptimizationOutputSupport:
category: str
description: str
pipe_name: str
removed_nodes: list[NodeRef] = field(default_factory=list)
added_nodes: list[NodeRef] = field(default_factory=list)
modified_nodes: list[NodeRef] = field(default_factory=list)
removed_initializers: list[str] = field(default_factory=list)
added_initializers: list[str] = field(default_factory=list)
modified_initializers: list[str] = field(default_factory=list)
operators: list[ProducedOperatorSupport] = field(default_factory=list)
error: str | None = None

Expand All @@ -105,6 +124,41 @@ def support_counts(self) -> dict[SupportLevel, int]:
"""Return a count of produced operators per support level."""
return dict(Counter(op.support for op in self.operators))

@staticmethod
def _node_ref_dict(ref: NodeRef) -> dict[str, object]:
"""Return the stable JSON representation of one graph-delta node."""
return {
"op_type": ref.op_type,
"name": ref.name,
"outputs": list(ref.outputs),
}

def to_dict(self) -> dict[str, object]:
"""Return actionable graph-delta and target-support evidence."""
data: dict[str, object] = {
"name": self.name,
"enable_flag": self.enable_flag,
"category": self.category,
"description": self.description,
"pipe_name": self.pipe_name,
"worst_support": self.worst_support.value,
"support_counts": {
level.value: count for level, count in self.support_counts().items()
},
"graph_delta": {
"removed_nodes": [self._node_ref_dict(ref) for ref in self.removed_nodes],
"added_nodes": [self._node_ref_dict(ref) for ref in self.added_nodes],
"modified_nodes": [self._node_ref_dict(ref) for ref in self.modified_nodes],
"removed_initializers": self.removed_initializers,
"added_initializers": self.added_initializers,
"modified_initializers": self.modified_initializers,
},
"operators": [operator.to_dict() for operator in self.operators],
}
if self.error is not None:
data["error"] = self.error
return data


def _produced_node_refs(
finding: CapabilityFinding,
Expand Down Expand Up @@ -137,6 +191,12 @@ def _check_one(
category=finding.category,
description=finding.description,
pipe_name=finding.pipe_name,
removed_nodes=list(finding.removed_nodes),
added_nodes=list(finding.added_nodes),
modified_nodes=list(finding.modified_nodes),
removed_initializers=list(finding.removed_initializers),
added_initializers=list(finding.added_initializers),
modified_initializers=list(finding.modified_initializers),
)

produced = _produced_node_refs(finding)
Expand Down
86 changes: 58 additions & 28 deletions src/winml/modelkit/commands/analyze.py
Original file line number Diff line number Diff line change
Expand Up @@ -1196,18 +1196,20 @@ def analyze(
# Optionally probe which optimizations would change the model and what
# operators they would introduce. This probe is target-independent, so
# materialize it once here and only re-run the (cheap) support lookup per
# EP/device below. Console-only feature — skipped in quiet mode.
# EP/device below. Quiet mode disables rendering, not data collection.
optim_outputs: list[tuple[Any, Any]] = []
if check_optim and not quiet:
optim_probe_error: str | None = None
if check_optim:
try:
import onnx

from ..optim import get_all_capabilities, iter_optimization_outputs

console.print(
"[dim]Probing optimization outputs "
"(this can take a while on large models)…[/dim]"
)
if not quiet:
console.print(
"[dim]Probing optimization outputs "
"(this can take a while on large models)…[/dim]"
)
optim_proto = onnx.load(str(model))
optim_outputs = list(iter_optimization_outputs(optim_proto, get_all_capabilities()))
# Each entry retains a full produced-model clone so the (cheap)
Expand All @@ -1221,6 +1223,7 @@ def analyze(
)
except Exception as exc:
logger.warning("Could not probe optimization outputs: %s", exc)
optim_probe_error = str(exc)
optim_outputs = []

# Model info header
Expand Down Expand Up @@ -1273,9 +1276,38 @@ def analyze(
ep_counter = 0
_no_data_eps: set[tuple[str, str]] = set() # EP/device pairs with no op rule data
analysis_results: list = []
optimization_support_payloads: list[dict[str, object] | None] = []
current_run_unknown_op = False
current_op_check_skipped = False

def _collect_optimization_support(
target_ep: EPName, target_device: str
) -> tuple[list[Any], dict[str, object]]:
"""Check one target and return renderable results plus JSON data."""
from ..analyze.optim_output import check_optimization_output_support

support_error: str | None = None
optim_support = []
if optim_probe_error is None:
try:
optim_support = check_optimization_output_support(
optim_outputs,
ep=target_ep,
device=target_device,
model_path=str(model),
)
except Exception as exc:
support_error = str(exc)
logger.warning("Could not check optimization output support: %s", exc)
payload: dict[str, object] = {
"ep_type": target_ep,
"device_type": target_device,
"probe_error": optim_probe_error,
"support_error": support_error,
"optimizations": [item.to_dict() for item in optim_support],
}
return optim_support, payload

def _current_ep_device_pair_display_name() -> str:
"""Return current EP/device display label, or empty when unset."""
if current_ep_device_pair is None:
Expand Down Expand Up @@ -1523,14 +1555,10 @@ def on_node_result(pattern_runtime: PatternRuntime) -> None:

# Optimization output support section (per-EP), when opted in.
if check_optim:
from ..analyze.optim_output import check_optimization_output_support

optim_support = check_optimization_output_support(
optim_outputs,
ep=target_ep,
device=target_device,
model_path=str(model),
optim_support, payload = _collect_optimization_support(
target_ep, target_device
)
optimization_support_payloads.append(payload)
_render_optim_output_support(
console,
optim_support,
Expand Down Expand Up @@ -1561,21 +1589,29 @@ def on_node_result(pattern_runtime: PatternRuntime) -> None:
)
analysis_results.append(result)

if check_optim:
_, payload = _collect_optimization_support(target_ep, target_device)
optimization_support_payloads.append(payload)

result = analysis_results[-1]

serialized_results: list[dict[str, object]] = []
json_mode = output_format == "json"
if output or json_mode:
for index, run_result in enumerate(analysis_results):
payload = json.loads(run_result.to_json())
if check_optim:
payload["optimization_output_support"] = optimization_support_payloads[index]
serialized_results.append(payload)

# Save JSON if requested
if output:
try:
output.parent.mkdir(parents=True, exist_ok=True)
if len(analysis_results) == 1:
output.write_text(result.to_json(), encoding="utf-8")
output.write_text(json.dumps(serialized_results[0], indent=2), encoding="utf-8")
else:
output.write_text(
json.dumps(
[json.loads(run_result.to_json()) for run_result in analysis_results]
),
encoding="utf-8",
)
output.write_text(json.dumps(serialized_results), encoding="utf-8")
logger.info("JSON results saved to: %s", output)
except OSError as e:
logger.error("Failed to write JSON output to %s: %s", output, e)
Expand Down Expand Up @@ -1654,17 +1690,11 @@ def on_node_result(pattern_runtime: PatternRuntime) -> None:
logger.debug("Config generation traceback:", exc_info=True)

# Emit JSON to stdout if requested
json_mode = output_format == "json"
if json_mode:
if len(analysis_results) == 1:
click.echo(result.to_json())
click.echo(json.dumps(serialized_results[0], indent=2))
else:
click.echo(
json.dumps(
[json.loads(run_result.to_json()) for run_result in analysis_results],
indent=2,
)
)
click.echo(json.dumps(serialized_results, indent=2))

# Exit code: 0 = fully supported, 1 = partial support
overall_supported = all(run_result.is_fully_supported() for run_result in analysis_results)
Expand Down
44 changes: 44 additions & 0 deletions tests/unit/analyze/test_optim_output.py
Original file line number Diff line number Diff line change
Expand Up @@ -164,6 +164,50 @@ def test_support_counts(self) -> None:
assert counts[SupportLevel.SUPPORTED] == 2
assert counts[SupportLevel.PARTIAL] == 1

def test_to_dict_includes_graph_delta_and_target_support(self) -> None:
"""Structured JSON retains actionable graph and support evidence."""
from winml.modelkit.optim import NodeRef

opt = OptimizationOutputSupport(
name="static-split-to-slice",
enable_flag="--enable-static-split-to-slice",
category="rewrite",
description="Replace static Split with Slice.",
pipe_name="algebraic",
removed_nodes=[NodeRef("Split", "split", ("a", "b"))],
added_nodes=[NodeRef("Slice", "slice_0", ("a",))],
modified_initializers=["starts"],
operators=[
ProducedOperatorSupport("Slice", "Slice 'slice_0'", "added", SupportLevel.SUPPORTED)
],
)

assert opt.to_dict() == {
"name": "static-split-to-slice",
"enable_flag": "--enable-static-split-to-slice",
"category": "rewrite",
"description": "Replace static Split with Slice.",
"pipe_name": "algebraic",
"worst_support": "supported",
"support_counts": {"supported": 1},
"graph_delta": {
"removed_nodes": [{"op_type": "Split", "name": "split", "outputs": ["a", "b"]}],
"added_nodes": [{"op_type": "Slice", "name": "slice_0", "outputs": ["a"]}],
"modified_nodes": [],
"removed_initializers": [],
"added_initializers": [],
"modified_initializers": ["starts"],
},
"operators": [
{
"op_type": "Slice",
"label": "Slice 'slice_0'",
"change": "added",
"support": "supported",
}
],
}


# =============================================================================
# END-TO-END SUPPORT CHECK
Expand Down
Loading
Loading