From d3ad8956774309a8e96d2d67c172a806450b7ca3 Mon Sep 17 00:00:00 2001 From: hualxie Date: Mon, 20 Jul 2026 16:20:08 +0800 Subject: [PATCH 1/3] feat(optim): add static Split-to-Slice rewrite --- .../modelkit/optim/capabilities/__init__.py | 2 + .../modelkit/optim/capabilities/algebraic.py | 21 + src/winml/modelkit/optim/pipes/__init__.py | 17 +- src/winml/modelkit/optim/pipes/algebraic.py | 423 ++++++++++++++++++ tests/unit/optim/pipes/test_pipe_algebraic.py | 267 +++++++++++ tests/unit/optim/test_optimizer.py | 6 +- 6 files changed, 733 insertions(+), 3 deletions(-) create mode 100644 src/winml/modelkit/optim/capabilities/algebraic.py create mode 100644 src/winml/modelkit/optim/pipes/algebraic.py create mode 100644 tests/unit/optim/pipes/test_pipe_algebraic.py diff --git a/src/winml/modelkit/optim/capabilities/__init__.py b/src/winml/modelkit/optim/capabilities/__init__.py index dcf954bf7..b219f1377 100644 --- a/src/winml/modelkit/optim/capabilities/__init__.py +++ b/src/winml/modelkit/optim/capabilities/__init__.py @@ -14,6 +14,7 @@ from . import ( activation, + algebraic, attention, conv, elimination, @@ -30,6 +31,7 @@ __all__ = [ "activation", + "algebraic", "attention", "conv", "elimination", diff --git a/src/winml/modelkit/optim/capabilities/algebraic.py b/src/winml/modelkit/optim/capabilities/algebraic.py new file mode 100644 index 000000000..d1bca5560 --- /dev/null +++ b/src/winml/modelkit/optim/capabilities/algebraic.py @@ -0,0 +1,21 @@ +# ------------------------------------------------------------------------- +# Copyright (c) Microsoft Corporation. All rights reserved. +# Licensed under the MIT License. +# -------------------------------------------------------------------------- +"""Opt-in, exact algebraic graph-rewrite capabilities.""" + +from __future__ import annotations + +from ..registry import BoolCapability, CapabilityCategory + + +STATIC_SPLIT_TO_SLICE = BoolCapability( + name="static-split-to-slice", + ort_name=None, + description=( + "Replace statically bounded Split operations with standard Slice operations " + "while preserving output tensors" + ), + category=CapabilityCategory.REWRITE, + default=False, +) diff --git a/src/winml/modelkit/optim/pipes/__init__.py b/src/winml/modelkit/optim/pipes/__init__.py index 6905a77b1..9315b71fa 100644 --- a/src/winml/modelkit/optim/pipes/__init__.py +++ b/src/winml/modelkit/optim/pipes/__init__.py @@ -11,6 +11,11 @@ from typing import Any +from .algebraic import ( + ALGEBRAIC_CAPABILITIES, + AlgebraicRewritePipe, + AlgebraicRewritePipeConfig, +) from .base import BasePipe, OptimizationError, PipeConfig, caps_dict from .fusion import ORTFusionPipe, ORTFusionPipeConfig from .graph import GRAPH_CAPABILITIES, ORTGraphPipe, ORTGraphPipeConfig @@ -22,12 +27,19 @@ # - ORTGraphPipe: ORT graph-level optimizations (C++ optimizer), including constant folding. # Runs first so downstream pipes see a constant-folded graph (e.g. Reshape shape inputs # become literal constants, enabling skeleton-based pattern matching). +# - AlgebraicRewritePipe: Exact topology-based algebraic rewrites (after ORT folding). # - RewritePipe: Pattern-based subgraph rewriting (runs after ORT constant folding so that # shape constants are visible, but before ORTFusionPipe so normalised patterns are # available for transformer fusions). # - ORTFusionPipe: ORT transformer fusions (Python optimizer) # - SurgeryPipe: Post-optimization model surgery (runs last to clamp constants after folding) -PIPES: list[type[BasePipe]] = [ORTGraphPipe, RewritePipe, ORTFusionPipe, SurgeryPipe] +PIPES: list[type[BasePipe]] = [ + ORTGraphPipe, + AlgebraicRewritePipe, + RewritePipe, + ORTFusionPipe, + SurgeryPipe, +] def get_all_capabilities() -> dict[str, Any]: @@ -43,9 +55,12 @@ def get_all_capabilities() -> dict[str, Any]: __all__ = [ + "ALGEBRAIC_CAPABILITIES", "GRAPH_CAPABILITIES", "PIPES", "SURGERY_CAPABILITIES", + "AlgebraicRewritePipe", + "AlgebraicRewritePipeConfig", "BasePipe", "ORTFusionPipe", "ORTFusionPipeConfig", diff --git a/src/winml/modelkit/optim/pipes/algebraic.py b/src/winml/modelkit/optim/pipes/algebraic.py new file mode 100644 index 000000000..a4f79b331 --- /dev/null +++ b/src/winml/modelkit/optim/pipes/algebraic.py @@ -0,0 +1,423 @@ +# ------------------------------------------------------------------------- +# Copyright (c) Microsoft Corporation. All rights reserved. +# Licensed under the MIT License. +# -------------------------------------------------------------------------- +"""Conservative, opt-in algebraic ONNX graph rewrites.""" + +from __future__ import annotations + +from dataclasses import dataclass +from typing import Any, ClassVar + +import numpy as np +import onnx + +from ..capabilities import algebraic +from .base import BasePipe, PipeConfig, caps_dict + + +ALGEBRAIC_CAPABILITIES: dict[str, Any] = caps_dict(algebraic.STATIC_SPLIT_TO_SLICE) + + +@dataclass +class AlgebraicRewritePipeConfig(PipeConfig): + """Configuration for exact algebraic rewrites.""" + + static_split_to_slice: bool = False + + +@dataclass +class _GraphIndex: + """Graph metadata required to identify statically bounded Split nodes.""" + + producers: dict[str, onnx.NodeProto] + consumers: dict[str, list[onnx.NodeProto]] + initializers: dict[str, onnx.TensorProto] + shapes: dict[str, tuple[int | None, ...]] + graph_outputs: set[str] + + @classmethod + def build(cls, model: onnx.ModelProto) -> _GraphIndex: + from onnx import numpy_helper + + graph = model.graph + producers: dict[str, onnx.NodeProto] = {} + consumers: dict[str, list[onnx.NodeProto]] = {} + for node in graph.node: + for output in node.output: + if output: + producers[output] = node + consumed_names = {input_name for input_name in node.input if input_name} + for attribute in node.attribute: + if attribute.type == onnx.AttributeProto.GRAPH: + consumed_names.update(_captured_tensor_names(attribute.g)) + elif attribute.type == onnx.AttributeProto.GRAPHS: + for nested_graph in attribute.graphs: + consumed_names.update(_captured_tensor_names(nested_graph)) + for input_name in consumed_names: + consumers.setdefault(input_name, []).append(node) + + initializers = {initializer.name: initializer for initializer in graph.initializer} + shapes: dict[str, tuple[int | None, ...]] = {} + for value_info in (*graph.input, *graph.value_info, *graph.output): + shape = _value_info_shape(value_info) + if shape is not None: + shapes[value_info.name] = shape + for name, initializer in initializers.items(): + shapes.setdefault(name, tuple(int(dim) for dim in initializer.dims)) + + for initializer in initializers.values(): + numpy_helper.to_array(initializer) + + return cls( + producers=producers, + consumers=consumers, + initializers=initializers, + shapes=shapes, + graph_outputs={output.name for output in graph.output if output.name}, + ) + + +class _NameAllocator: + """Allocate names without relying on optional or duplicated node names.""" + + def __init__(self, model: onnx.ModelProto) -> None: + graph = model.graph + self._used = { + name + for name in ( + [initializer.name for initializer in graph.initializer] + + [value.name for value in graph.input] + + [value.name for value in graph.value_info] + + [value.name for value in graph.output] + + [node.name for node in graph.node] + + [output for node in graph.node for output in node.output] + ) + if name + } + + def new(self, prefix: str) -> str: + candidate = prefix + suffix = 0 + while candidate in self._used: + suffix += 1 + candidate = f"{prefix}_{suffix}" + self._used.add(candidate) + return candidate + + +def _value_info_shape(value_info: onnx.ValueInfoProto) -> tuple[int | None, ...] | None: + """Return a tensor shape, preserving unknown dimensions as ``None``.""" + if not value_info.type.HasField("tensor_type"): + return None + tensor_type = value_info.type.tensor_type + if not tensor_type.HasField("shape"): + return None + dimensions = [ + int(dimension.dim_value) if dimension.HasField("dim_value") else None + for dimension in tensor_type.shape.dim + ] + return tuple(dimensions) + + +def _attribute(node: onnx.NodeProto, name: str, default: Any = None) -> Any: + from onnx import helper + + for attribute in node.attribute: + if attribute.name == name: + return helper.get_attribute_value(attribute) + return default + + +def _constant_array(index: _GraphIndex, name: str) -> np.ndarray | None: + """Read an initializer or a regular ONNX Constant value.""" + from onnx import numpy_helper + + if not name: + return None + initializer = index.initializers.get(name) + if initializer is not None: + return np.asarray(numpy_helper.to_array(initializer)) + + producer = index.producers.get(name) + if producer is None or producer.op_type != "Constant": + return None + value = _attribute(producer, "value") + if value is not None: + try: + return np.asarray(numpy_helper.to_array(value)) + except (TypeError, ValueError): + return None + for attribute_name in ("value_float", "value_floats", "value_int", "value_ints"): + attribute_value = _attribute(producer, attribute_name) + if attribute_value is not None: + return np.asarray(attribute_value) + return None + + +def _constant_ints(index: _GraphIndex, name: str) -> list[int] | None: + values = _constant_array(index, name) + if values is None or not np.issubdtype(values.dtype, np.integer): + return None + return [int(value) for value in values.reshape(-1).tolist()] + + +def _single_attribute_or_input_ints( + index: _GraphIndex, + node: onnx.NodeProto, + attribute_name: str, + input_index: int | None, +) -> tuple[list[int] | None, bool]: + """Read legacy attributes and newer constant inputs.""" + attribute_value = _attribute(node, attribute_name) + from_attribute = None + if attribute_value is not None: + try: + values = np.asarray(attribute_value) + if not np.issubdtype(values.dtype, np.integer): + return None, True + from_attribute = [int(value) for value in values.reshape(-1).tolist()] + except (TypeError, ValueError): + return None, True + + from_input = None + if input_index is not None and len(node.input) > input_index and node.input[input_index]: + from_input = _constant_ints(index, node.input[input_index]) + if from_input is None: + return None, True + + if from_attribute is not None and from_input is not None and from_attribute != from_input: + return None, True + return from_input if from_input is not None else from_attribute, False + + +def _new_initializer( + model: onnx.ModelProto, + allocator: _NameAllocator, + values: np.ndarray, + prefix: str, +) -> str: + from onnx import numpy_helper + + name = allocator.new(prefix) + model.graph.initializer.append(numpy_helper.from_array(np.asarray(values), name)) + return name + + +def _remove_nodes(model: onnx.ModelProto, nodes: set[int]) -> None: + remaining = [node for node in model.graph.node if id(node) not in nodes] + del model.graph.node[:] + model.graph.node.extend(remaining) + + +def _captured_tensor_names(graph: onnx.GraphProto) -> set[str]: + """Return names a nested graph resolves from an enclosing scope.""" + locally_defined = {value.name for value in graph.input if value.name} + locally_defined.update(initializer.name for initializer in graph.initializer) + locally_defined.update(output for node in graph.node for output in node.output if output) + referenced = {input_name for node in graph.node for input_name in node.input if input_name} + referenced.update(output.name for output in graph.output if output.name) + for node in graph.node: + for attribute in node.attribute: + if attribute.type == onnx.AttributeProto.GRAPH: + referenced.update(_captured_tensor_names(attribute.g)) + elif attribute.type == onnx.AttributeProto.GRAPHS: + for nested_graph in attribute.graphs: + referenced.update(_captured_tensor_names(nested_graph)) + return referenced - locally_defined + + +def _referenced_tensor_names(graph: onnx.GraphProto) -> set[str]: + """Collect tensor names referenced by a graph or its nested subgraphs.""" + referenced = {value.name for value in (*graph.input, *graph.output) if value.name} + for node in graph.node: + referenced.update(input_name for input_name in node.input if input_name) + for attribute in node.attribute: + if attribute.type == onnx.AttributeProto.GRAPH: + referenced.update(_referenced_tensor_names(attribute.g)) + elif attribute.type == onnx.AttributeProto.GRAPHS: + for nested_graph in attribute.graphs: + referenced.update(_referenced_tensor_names(nested_graph)) + return referenced + + +def _prune_unused_initializers(model: onnx.ModelProto) -> None: + used = _referenced_tensor_names(model.graph) + remaining = [initializer for initializer in model.graph.initializer if initializer.name in used] + del model.graph.initializer[:] + model.graph.initializer.extend(remaining) + + +def _prune_generated_slices(model: onnx.ModelProto, introduced: set[str]) -> None: + """Remove generated Slice nodes whose outputs are entirely dead.""" + while True: + index = _GraphIndex.build(model) + removable = { + id(node) + for node in model.graph.node + if node.name in introduced + and all( + output and output not in index.graph_outputs and not index.consumers.get(output) + for output in node.output + ) + } + if not removable: + return + _remove_nodes(model, removable) + + +def _prune_dead_constant_nodes(model: onnx.ModelProto) -> None: + while True: + index = _GraphIndex.build(model) + removable = { + id(node) + for node in model.graph.node + if node.op_type == "Constant" + and all( + output and output not in index.graph_outputs and not index.consumers.get(output) + for output in node.output + ) + } + if not removable: + return + _remove_nodes(model, removable) + + +def _split_boundaries( + index: _GraphIndex, + node: onnx.NodeProto, + input_name: str, +) -> tuple[int, list[tuple[int, int]]] | None: + """Return a static Split axis and output boundaries.""" + input_shape = index.shapes.get(input_name) + if input_shape is None or len(node.output) == 0: + return None + + axis_input_index = 2 if len(node.input) > 2 else None + axis_values, axis_conflict = _single_attribute_or_input_ints( + index, node, "axis", axis_input_index + ) + if axis_conflict or (axis_values is not None and len(axis_values) != 1): + return None + axis = axis_values[0] if axis_values is not None else 0 + if axis < -len(input_shape) or axis >= len(input_shape): + return None + axis %= len(input_shape) + axis_size = input_shape[axis] + if axis_size is None: + return None + + split_values, split_conflict = _single_attribute_or_input_ints(index, node, "split", 1) + if split_conflict: + return None + if split_values is None: + if axis_size <= 0 or axis_size % len(node.output) != 0: + return None + split_values = [axis_size // len(node.output)] * len(node.output) + if len(split_values) != len(node.output) or any(value <= 0 for value in split_values): + return None + if sum(split_values) != axis_size: + return None + + boundaries: list[tuple[int, int]] = [] + start = 0 + for size in split_values: + boundaries.append((start, start + size)) + start += size + return axis, boundaries + + +def _rewrite_static_splits( + model: onnx.ModelProto, + allocator: _NameAllocator, + introduced_nodes: set[str], +) -> None: + """Replace statically bounded Split nodes with input-form Slice nodes.""" + from onnx import helper + + index = _GraphIndex.build(model) + opset = next( + (int(opset.version) for opset in model.opset_import if opset.domain in ("", "ai.onnx")), + 0, + ) + if opset and opset < 10: + return + + replacements: dict[int, list[onnx.NodeProto]] = {} + for split in list(model.graph.node): + if split.op_type != "Split" or len(split.input) < 1 or not split.input[0]: + continue + if any(not output for output in split.output): + continue + info = _split_boundaries(index, split, split.input[0]) + if info is None: + continue + axis, boundaries = info + replacement: list[onnx.NodeProto] = [] + for output_index, (start, end) in enumerate(boundaries): + starts_name = _new_initializer( + model, allocator, np.asarray([start], dtype=np.int64), "algebraic_slice_starts" + ) + ends_name = _new_initializer( + model, allocator, np.asarray([end], dtype=np.int64), "algebraic_slice_ends" + ) + axes_name = _new_initializer( + model, allocator, np.asarray([axis], dtype=np.int64), "algebraic_slice_axes" + ) + steps_name = _new_initializer( + model, allocator, np.asarray([1], dtype=np.int64), "algebraic_slice_steps" + ) + replacement_node = helper.make_node( + "Slice", + [split.input[0], starts_name, ends_name, axes_name, steps_name], + [split.output[output_index]], + name=allocator.new("algebraic_split_slice"), + ) + replacement.append(replacement_node) + introduced_nodes.add(replacement_node.name) + replacements[id(split)] = replacement + + if not replacements: + return + rewritten: list[onnx.NodeProto] = [] + for node in model.graph.node: + rewritten.extend(replacements.get(id(node), [node])) + del model.graph.node[:] + model.graph.node.extend(rewritten) + + +class AlgebraicRewritePipe(BasePipe[AlgebraicRewritePipeConfig]): + """Replace statically bounded Split nodes with Slice nodes.""" + + name: ClassVar[str] = "algebraic_rewrite" + capabilities: ClassVar[dict[str, Any]] = ALGEBRAIC_CAPABILITIES + + @classmethod + def build_config(cls, **kwargs: Any) -> AlgebraicRewritePipeConfig: + """Build the static Split-to-Slice configuration.""" + return AlgebraicRewritePipeConfig( + static_split_to_slice=kwargs.get("static_split_to_slice", False) + ) + + @classmethod + def should_process(cls, config: AlgebraicRewritePipeConfig) -> bool: + """Return whether static Split-to-Slice rewriting is enabled.""" + return config.static_split_to_slice + + def process( + self, + model: onnx.ModelProto, + config: AlgebraicRewritePipeConfig, + ) -> onnx.ModelProto: + """Rewrite eligible static Split nodes in a copy of the model.""" + if not self.should_process(config): + return model + + result = onnx.ModelProto() + result.CopyFrom(model) + introduced_nodes: set[str] = set() + _rewrite_static_splits(result, _NameAllocator(result), introduced_nodes) + _prune_generated_slices(result, introduced_nodes) + _prune_dead_constant_nodes(result) + _prune_unused_initializers(result) + return result diff --git a/tests/unit/optim/pipes/test_pipe_algebraic.py b/tests/unit/optim/pipes/test_pipe_algebraic.py new file mode 100644 index 000000000..749fba3c9 --- /dev/null +++ b/tests/unit/optim/pipes/test_pipe_algebraic.py @@ -0,0 +1,267 @@ +# ------------------------------------------------------------------------- +# Copyright (c) Microsoft Corporation. All rights reserved. +# Licensed under the MIT License. +# -------------------------------------------------------------------------- +"""Generated-graph tests for static Split-to-Slice rewriting.""" + +from __future__ import annotations + +from typing import TYPE_CHECKING + +import numpy as np +import onnx +import onnxruntime as ort +from click.testing import CliRunner +from onnx import TensorProto, helper, numpy_helper + +from winml.modelkit.commands.optimize import optimize +from winml.modelkit.optim import get_all_capabilities, optimize_onnx +from winml.modelkit.optim.pipes import ( + PIPES, + AlgebraicRewritePipe, + AlgebraicRewritePipeConfig, +) + + +if TYPE_CHECKING: + from collections.abc import Sequence + + +def _tensor(name: str, values: np.ndarray) -> onnx.TensorProto: + return numpy_helper.from_array(np.asarray(values), name) + + +def _model( + nodes: Sequence[onnx.NodeProto], + inputs: Sequence[onnx.ValueInfoProto], + outputs: Sequence[onnx.ValueInfoProto], + initializers: Sequence[onnx.TensorProto], + value_info: Sequence[onnx.ValueInfoProto] = (), +) -> onnx.ModelProto: + graph = helper.make_graph( + list(nodes), + "generated_algebraic_graph", + list(inputs), + list(outputs), + initializer=list(initializers), + value_info=list(value_info), + ) + model = helper.make_model(graph, opset_imports=[helper.make_opsetid("", 17)]) + model.ir_version = 8 + return model + + +def _info(name: str, shape: Sequence[int | None]) -> onnx.ValueInfoProto: + return helper.make_tensor_value_info(name, TensorProto.FLOAT, list(shape)) + + +def _run(model: onnx.ModelProto, values: dict[str, np.ndarray]) -> list[np.ndarray]: + session = ort.InferenceSession( + model.SerializeToString(), + providers=["CPUExecutionProvider"], + ) + return session.run(None, values) + + +def _assert_valid_with_inferred_shapes(model: onnx.ModelProto) -> None: + onnx.checker.check_model(model) + inferred = onnx.shape_inference.infer_shapes(model) + assert len(inferred.graph.output) == len(model.graph.output) + + +class TestAlgebraicRegistration: + """Verify capability registration, flags, and pipe ordering.""" + + def test_capability_is_opt_in(self) -> None: + capability = get_all_capabilities()["static-split-to-slice"] + assert capability.default is False + assert capability.cli_flags() == ( + "--enable-static-split-to-slice", + "--disable-static-split-to-slice", + ) + + config = AlgebraicRewritePipe.build_config(static_split_to_slice=True) + assert config.static_split_to_slice is True + + def test_cli_lists_algebraic_flag(self) -> None: + result = CliRunner().invoke(optimize, ["--list-capabilities"]) + assert result.exit_code == 0 + assert "--enable-static-split-to-slice" in result.output + + def test_pipe_is_after_ort_graph_and_before_cleanup(self) -> None: + names = [pipe.name for pipe in PIPES] + assert names.index("ort_graph") < names.index("algebraic_rewrite") + assert names.index("algebraic_rewrite") < names.index("surgery") + assert PIPES[names.index("algebraic_rewrite")] is AlgebraicRewritePipe + assert not AlgebraicRewritePipe.should_process(AlgebraicRewritePipeConfig()) + + +class TestStaticSplitToSlice: + """Test static Split replacement using generated data.""" + + def test_positive_equivalence_and_preserved_outputs(self) -> None: + rng = np.random.default_rng(10) + x = helper.make_tensor_value_info("x", TensorProto.FLOAT, [1, 6, 2]) + outputs = [_info("left", [1, 2, 2]), _info("right", [1, 4, 2])] + split = helper.make_node( + "Split", + ["x", "split_sizes"], + ["left", "right"], + name="", + axis=1, + ) + model = _model( + [split], + [x], + outputs, + [_tensor("split_sizes", np.asarray([2, 4], dtype=np.int64))], + ) + transformed = AlgebraicRewritePipe().process( + model, + AlgebraicRewritePipeConfig(static_split_to_slice=True), + ) + + assert [node.op_type for node in transformed.graph.node] == ["Slice", "Slice"] + assert [node.output[0] for node in transformed.graph.node] == ["left", "right"] + assert [output.name for output in transformed.graph.output] == ["left", "right"] + assert "split_sizes" not in { + initializer.name for initializer in transformed.graph.initializer + } + _assert_valid_with_inferred_shapes(transformed) + values = {"x": rng.normal(size=(1, 6, 2)).astype(np.float32)} + for original, rewritten in zip(_run(model, values), _run(transformed, values), strict=True): + np.testing.assert_allclose(original, rewritten, rtol=0, atol=0) + + def test_equal_split_and_name_collisions_are_safe(self) -> None: + x = helper.make_tensor_value_info("x", TensorProto.FLOAT, [1, 4, 2]) + keep = _info("keep", [1, 2, 2]) + split = helper.make_node( + "Split", + ["x"], + ["part_a", "part_b"], + name="", + axis=1, + ) + identity = helper.make_node( + "Identity", + ["part_b"], + ["keep"], + name="algebraic_split_slice", + ) + model = _model( + [split, identity], + [x], + [keep, _info("part_a", [1, 2, 2])], + [], + value_info=[_info("part_a", [1, 2, 2]), _info("part_b", [1, 2, 2])], + ) + transformed = AlgebraicRewritePipe().process( + model, + AlgebraicRewritePipeConfig(static_split_to_slice=True), + ) + generated = [node for node in transformed.graph.node if node.op_type == "Slice"] + assert len(generated) == 2 + assert len({node.name for node in transformed.graph.node}) == len(transformed.graph.node) + assert all(node.name for node in generated) + assert {node.output[0] for node in generated} == {"part_a", "part_b"} + _assert_valid_with_inferred_shapes(transformed) + + def test_dynamic_equal_split_and_malformed_split_are_unchanged(self) -> None: + x = helper.make_tensor_value_info("x", TensorProto.FLOAT, [1, None, 2]) + dynamic_equal = helper.make_node("Split", ["x"], ["a", "b"], axis=1) + malformed = helper.make_node( + "Split", + ["x", "bad_sizes"], + ["c", "d"], + axis=1, + ) + model = _model( + [dynamic_equal, malformed], + [x], + [_info("a", [1, None, 2]), _info("b", [1, None, 2])], + [_tensor("bad_sizes", np.asarray([1, 1], dtype=np.int64))], + ) + transformed = AlgebraicRewritePipe().process( + model, + AlgebraicRewritePipeConfig(static_split_to_slice=True), + ) + assert [node.op_type for node in transformed.graph.node] == ["Split", "Split"] + + def test_dead_generated_slice_and_constants_are_pruned(self) -> None: + x = helper.make_tensor_value_info("x", TensorProto.FLOAT, [1, 4, 2]) + model = _model( + [helper.make_node("Split", ["x"], ["left", "unused"], axis=1)], + [x], + [_info("left", [1, 2, 2])], + [], + ) + transformed = AlgebraicRewritePipe().process( + model, + AlgebraicRewritePipeConfig(static_split_to_slice=True), + ) + assert [node.op_type for node in transformed.graph.node] == ["Slice"] + assert transformed.graph.node[0].output[0] == "left" + assert len(transformed.graph.initializer) == 4 + + def test_nested_subgraph_captures_keep_generated_slices_live(self) -> None: + x = helper.make_tensor_value_info("x", TensorProto.FLOAT, [1, 4, 2]) + then_branch = helper.make_graph( + [helper.make_node("Identity", ["left"], ["then_output"])], + "then_branch", + [], + [_info("then_output", [1, 2, 2])], + ) + else_branch = helper.make_graph( + [helper.make_node("Identity", ["right"], ["else_output"])], + "else_branch", + [], + [_info("else_output", [1, 2, 2])], + ) + model = _model( + [ + helper.make_node("Split", ["x"], ["left", "right"], axis=1), + helper.make_node( + "If", + ["condition"], + ["y"], + then_branch=then_branch, + else_branch=else_branch, + ), + ], + [x], + [_info("y", [1, 2, 2])], + [_tensor("condition", np.asarray(True, dtype=np.bool_))], + value_info=[_info("left", [1, 2, 2]), _info("right", [1, 2, 2])], + ) + values = {"x": np.arange(8, dtype=np.float32).reshape(1, 4, 2)} + transformed = AlgebraicRewritePipe().process( + model, + AlgebraicRewritePipeConfig(static_split_to_slice=True), + ) + assert [node.op_type for node in transformed.graph.node] == [ + "Slice", + "Slice", + "If", + ] + _assert_valid_with_inferred_shapes(transformed) + np.testing.assert_array_equal(_run(model, values), _run(transformed, values)) + + def test_public_optimize_path_is_idempotent(self) -> None: + x = helper.make_tensor_value_info("x", TensorProto.FLOAT, [1, 4, 2]) + model = _model( + [helper.make_node("Split", ["x"], ["a", "b"], name="", axis=1)], + [x], + [_info("a", [1, 2, 2]), _info("b", [1, 2, 2])], + [], + ) + transformed = optimize_onnx(model, static_split_to_slice=True) + second = optimize_onnx(transformed, static_split_to_slice=True) + assert all(node.op_type != "Split" for node in transformed.graph.node) + assert [ + (node.op_type, tuple(node.input), tuple(node.output), node.name) + for node in transformed.graph.node + ] == [ + (node.op_type, tuple(node.input), tuple(node.output), node.name) + for node in second.graph.node + ] + _assert_valid_with_inferred_shapes(second) diff --git a/tests/unit/optim/test_optimizer.py b/tests/unit/optim/test_optimizer.py index 1e51af2b7..3095e30a7 100644 --- a/tests/unit/optim/test_optimizer.py +++ b/tests/unit/optim/test_optimizer.py @@ -698,14 +698,16 @@ def test_resolve_dependencies_method(self) -> None: def test_registered_pipes_count(self) -> None: """Verify the expected number of pipes are registered.""" Optimizer._initialize_pipes() - # Currently: ORTGraphPipe, RewritePipe, ORTFusionPipe, SurgeryPipe - assert len(Optimizer.pipes) == 4 + # Currently: ORTGraphPipe, AlgebraicRewritePipe, RewritePipe, + # ORTFusionPipe, SurgeryPipe + assert len(Optimizer.pipes) == 5 def test_registered_pipe_names(self) -> None: """Verify expected pipe names are registered.""" Optimizer._initialize_pipes() names = {pipe_class.name for pipe_class in Optimizer.pipes} assert "rewrite" in names + assert "algebraic_rewrite" in names assert "ort_graph" in names assert "ort_fusion" in names assert "surgery" in names From 8d4144af1182bf38016fa98d6d2a808ffab8a8b5 Mon Sep 17 00:00:00 2001 From: Hualiang Xie Date: Mon, 20 Jul 2026 16:48:16 +0800 Subject: [PATCH 2/3] fix(optim): use consistent ONNX imports --- src/winml/modelkit/optim/pipes/algebraic.py | 22 +++------ tests/unit/optim/pipes/test_pipe_algebraic.py | 47 +++++++++---------- 2 files changed, 29 insertions(+), 40 deletions(-) diff --git a/src/winml/modelkit/optim/pipes/algebraic.py b/src/winml/modelkit/optim/pipes/algebraic.py index a4f79b331..eec286917 100644 --- a/src/winml/modelkit/optim/pipes/algebraic.py +++ b/src/winml/modelkit/optim/pipes/algebraic.py @@ -38,8 +38,6 @@ class _GraphIndex: @classmethod def build(cls, model: onnx.ModelProto) -> _GraphIndex: - from onnx import numpy_helper - graph = model.graph producers: dict[str, onnx.NodeProto] = {} consumers: dict[str, list[onnx.NodeProto]] = {} @@ -67,7 +65,7 @@ def build(cls, model: onnx.ModelProto) -> _GraphIndex: shapes.setdefault(name, tuple(int(dim) for dim in initializer.dims)) for initializer in initializers.values(): - numpy_helper.to_array(initializer) + onnx.numpy_helper.to_array(initializer) return cls( producers=producers, @@ -121,23 +119,19 @@ def _value_info_shape(value_info: onnx.ValueInfoProto) -> tuple[int | None, ...] def _attribute(node: onnx.NodeProto, name: str, default: Any = None) -> Any: - from onnx import helper - for attribute in node.attribute: if attribute.name == name: - return helper.get_attribute_value(attribute) + return onnx.helper.get_attribute_value(attribute) return default def _constant_array(index: _GraphIndex, name: str) -> np.ndarray | None: """Read an initializer or a regular ONNX Constant value.""" - from onnx import numpy_helper - if not name: return None initializer = index.initializers.get(name) if initializer is not None: - return np.asarray(numpy_helper.to_array(initializer)) + return np.asarray(onnx.numpy_helper.to_array(initializer)) producer = index.producers.get(name) if producer is None or producer.op_type != "Constant": @@ -145,7 +139,7 @@ def _constant_array(index: _GraphIndex, name: str) -> np.ndarray | None: value = _attribute(producer, "value") if value is not None: try: - return np.asarray(numpy_helper.to_array(value)) + return np.asarray(onnx.numpy_helper.to_array(value)) except (TypeError, ValueError): return None for attribute_name in ("value_float", "value_floats", "value_int", "value_ints"): @@ -197,10 +191,8 @@ def _new_initializer( values: np.ndarray, prefix: str, ) -> str: - from onnx import numpy_helper - name = allocator.new(prefix) - model.graph.initializer.append(numpy_helper.from_array(np.asarray(values), name)) + model.graph.initializer.append(onnx.numpy_helper.from_array(np.asarray(values), name)) return name @@ -333,8 +325,6 @@ def _rewrite_static_splits( introduced_nodes: set[str], ) -> None: """Replace statically bounded Split nodes with input-form Slice nodes.""" - from onnx import helper - index = _GraphIndex.build(model) opset = next( (int(opset.version) for opset in model.opset_import if opset.domain in ("", "ai.onnx")), @@ -367,7 +357,7 @@ def _rewrite_static_splits( steps_name = _new_initializer( model, allocator, np.asarray([1], dtype=np.int64), "algebraic_slice_steps" ) - replacement_node = helper.make_node( + replacement_node = onnx.helper.make_node( "Slice", [split.input[0], starts_name, ends_name, axes_name, steps_name], [split.output[output_index]], diff --git a/tests/unit/optim/pipes/test_pipe_algebraic.py b/tests/unit/optim/pipes/test_pipe_algebraic.py index 749fba3c9..7031e2449 100644 --- a/tests/unit/optim/pipes/test_pipe_algebraic.py +++ b/tests/unit/optim/pipes/test_pipe_algebraic.py @@ -12,7 +12,6 @@ import onnx import onnxruntime as ort from click.testing import CliRunner -from onnx import TensorProto, helper, numpy_helper from winml.modelkit.commands.optimize import optimize from winml.modelkit.optim import get_all_capabilities, optimize_onnx @@ -28,7 +27,7 @@ def _tensor(name: str, values: np.ndarray) -> onnx.TensorProto: - return numpy_helper.from_array(np.asarray(values), name) + return onnx.numpy_helper.from_array(np.asarray(values), name) def _model( @@ -38,7 +37,7 @@ def _model( initializers: Sequence[onnx.TensorProto], value_info: Sequence[onnx.ValueInfoProto] = (), ) -> onnx.ModelProto: - graph = helper.make_graph( + graph = onnx.helper.make_graph( list(nodes), "generated_algebraic_graph", list(inputs), @@ -46,13 +45,13 @@ def _model( initializer=list(initializers), value_info=list(value_info), ) - model = helper.make_model(graph, opset_imports=[helper.make_opsetid("", 17)]) + model = onnx.helper.make_model(graph, opset_imports=[onnx.helper.make_opsetid("", 17)]) model.ir_version = 8 return model def _info(name: str, shape: Sequence[int | None]) -> onnx.ValueInfoProto: - return helper.make_tensor_value_info(name, TensorProto.FLOAT, list(shape)) + return onnx.helper.make_tensor_value_info(name, onnx.TensorProto.FLOAT, list(shape)) def _run(model: onnx.ModelProto, values: dict[str, np.ndarray]) -> list[np.ndarray]: @@ -101,9 +100,9 @@ class TestStaticSplitToSlice: def test_positive_equivalence_and_preserved_outputs(self) -> None: rng = np.random.default_rng(10) - x = helper.make_tensor_value_info("x", TensorProto.FLOAT, [1, 6, 2]) + x = onnx.helper.make_tensor_value_info("x", onnx.TensorProto.FLOAT, [1, 6, 2]) outputs = [_info("left", [1, 2, 2]), _info("right", [1, 4, 2])] - split = helper.make_node( + split = onnx.helper.make_node( "Split", ["x", "split_sizes"], ["left", "right"], @@ -133,16 +132,16 @@ def test_positive_equivalence_and_preserved_outputs(self) -> None: np.testing.assert_allclose(original, rewritten, rtol=0, atol=0) def test_equal_split_and_name_collisions_are_safe(self) -> None: - x = helper.make_tensor_value_info("x", TensorProto.FLOAT, [1, 4, 2]) + x = onnx.helper.make_tensor_value_info("x", onnx.TensorProto.FLOAT, [1, 4, 2]) keep = _info("keep", [1, 2, 2]) - split = helper.make_node( + split = onnx.helper.make_node( "Split", ["x"], ["part_a", "part_b"], name="", axis=1, ) - identity = helper.make_node( + identity = onnx.helper.make_node( "Identity", ["part_b"], ["keep"], @@ -167,9 +166,9 @@ def test_equal_split_and_name_collisions_are_safe(self) -> None: _assert_valid_with_inferred_shapes(transformed) def test_dynamic_equal_split_and_malformed_split_are_unchanged(self) -> None: - x = helper.make_tensor_value_info("x", TensorProto.FLOAT, [1, None, 2]) - dynamic_equal = helper.make_node("Split", ["x"], ["a", "b"], axis=1) - malformed = helper.make_node( + x = onnx.helper.make_tensor_value_info("x", onnx.TensorProto.FLOAT, [1, None, 2]) + dynamic_equal = onnx.helper.make_node("Split", ["x"], ["a", "b"], axis=1) + malformed = onnx.helper.make_node( "Split", ["x", "bad_sizes"], ["c", "d"], @@ -188,9 +187,9 @@ def test_dynamic_equal_split_and_malformed_split_are_unchanged(self) -> None: assert [node.op_type for node in transformed.graph.node] == ["Split", "Split"] def test_dead_generated_slice_and_constants_are_pruned(self) -> None: - x = helper.make_tensor_value_info("x", TensorProto.FLOAT, [1, 4, 2]) + x = onnx.helper.make_tensor_value_info("x", onnx.TensorProto.FLOAT, [1, 4, 2]) model = _model( - [helper.make_node("Split", ["x"], ["left", "unused"], axis=1)], + [onnx.helper.make_node("Split", ["x"], ["left", "unused"], axis=1)], [x], [_info("left", [1, 2, 2])], [], @@ -204,23 +203,23 @@ def test_dead_generated_slice_and_constants_are_pruned(self) -> None: assert len(transformed.graph.initializer) == 4 def test_nested_subgraph_captures_keep_generated_slices_live(self) -> None: - x = helper.make_tensor_value_info("x", TensorProto.FLOAT, [1, 4, 2]) - then_branch = helper.make_graph( - [helper.make_node("Identity", ["left"], ["then_output"])], + x = onnx.helper.make_tensor_value_info("x", onnx.TensorProto.FLOAT, [1, 4, 2]) + then_branch = onnx.helper.make_graph( + [onnx.helper.make_node("Identity", ["left"], ["then_output"])], "then_branch", [], [_info("then_output", [1, 2, 2])], ) - else_branch = helper.make_graph( - [helper.make_node("Identity", ["right"], ["else_output"])], + else_branch = onnx.helper.make_graph( + [onnx.helper.make_node("Identity", ["right"], ["else_output"])], "else_branch", [], [_info("else_output", [1, 2, 2])], ) model = _model( [ - helper.make_node("Split", ["x"], ["left", "right"], axis=1), - helper.make_node( + onnx.helper.make_node("Split", ["x"], ["left", "right"], axis=1), + onnx.helper.make_node( "If", ["condition"], ["y"], @@ -247,9 +246,9 @@ def test_nested_subgraph_captures_keep_generated_slices_live(self) -> None: np.testing.assert_array_equal(_run(model, values), _run(transformed, values)) def test_public_optimize_path_is_idempotent(self) -> None: - x = helper.make_tensor_value_info("x", TensorProto.FLOAT, [1, 4, 2]) + x = onnx.helper.make_tensor_value_info("x", onnx.TensorProto.FLOAT, [1, 4, 2]) model = _model( - [helper.make_node("Split", ["x"], ["a", "b"], name="", axis=1)], + [onnx.helper.make_node("Split", ["x"], ["a", "b"], name="", axis=1)], [x], [_info("a", [1, 2, 2]), _info("b", [1, 2, 2])], [], From ae9f4ed445545e82ebc9ca66a2c3d22d00a32bc0 Mon Sep 17 00:00:00 2001 From: Hualiang Xie Date: Tue, 21 Jul 2026 14:06:47 +0800 Subject: [PATCH 3/3] fix(optim): streamline Split graph analysis --- src/winml/modelkit/optim/pipes/algebraic.py | 12 +++--------- 1 file changed, 3 insertions(+), 9 deletions(-) diff --git a/src/winml/modelkit/optim/pipes/algebraic.py b/src/winml/modelkit/optim/pipes/algebraic.py index eec286917..3de2d64d9 100644 --- a/src/winml/modelkit/optim/pipes/algebraic.py +++ b/src/winml/modelkit/optim/pipes/algebraic.py @@ -64,9 +64,6 @@ def build(cls, model: onnx.ModelProto) -> _GraphIndex: for name, initializer in initializers.items(): shapes.setdefault(name, tuple(int(dim) for dim in initializer.dims)) - for initializer in initializers.values(): - onnx.numpy_helper.to_array(initializer) - return cls( producers=producers, consumers=consumers, @@ -285,13 +282,10 @@ def _split_boundaries( if input_shape is None or len(node.output) == 0: return None - axis_input_index = 2 if len(node.input) > 2 else None - axis_values, axis_conflict = _single_attribute_or_input_ints( - index, node, "axis", axis_input_index - ) - if axis_conflict or (axis_values is not None and len(axis_values) != 1): + axis_value = _attribute(node, "axis", 0) + if not isinstance(axis_value, (int, np.integer)): return None - axis = axis_values[0] if axis_values is not None else 0 + axis = int(axis_value) if axis < -len(input_shape) or axis >= len(input_shape): return None axis %= len(input_shape)