diff --git a/docs/tutorial/rewriter/examples/broadcast_matmul.py b/docs/tutorial/rewriter/examples/broadcast_matmul.py index cf56b49f07..73d8c45a2f 100644 --- a/docs/tutorial/rewriter/examples/broadcast_matmul.py +++ b/docs/tutorial/rewriter/examples/broadcast_matmul.py @@ -12,9 +12,10 @@ import logging import onnx +import onnx_ir as ir import onnxscript -from onnxscript import FLOAT, ir, opset18, script +from onnxscript import FLOAT, opset18, script from onnxscript.rewriter import pattern logger = logging.getLogger(__name__) diff --git a/docs/tutorial/rewriter/examples/erfgelu.py b/docs/tutorial/rewriter/examples/erfgelu.py index e042d9f337..ef361117e5 100644 --- a/docs/tutorial/rewriter/examples/erfgelu.py +++ b/docs/tutorial/rewriter/examples/erfgelu.py @@ -11,9 +11,10 @@ import math import onnx +import onnx_ir as ir import onnxscript -from onnxscript import FLOAT, ir, opset18, script +from onnxscript import FLOAT, opset18, script from onnxscript.rewriter import pattern diff --git a/examples/pattern_matching_example.py b/examples/pattern_matching_example.py index 8de09ecd6a..6cf9b22cf7 100644 --- a/examples/pattern_matching_example.py +++ b/examples/pattern_matching_example.py @@ -3,8 +3,8 @@ """Example demonstrating the new pattern matching functionality.""" import onnx.parser +import onnx_ir as ir -from onnxscript import ir from onnxscript.rewriter import pattern diff --git a/examples/pattern_rewriting.py b/examples/pattern_rewriting.py index fd84d7f3cb..9433f6c647 100644 --- a/examples/pattern_rewriting.py +++ b/examples/pattern_rewriting.py @@ -14,8 +14,8 @@ import onnx import onnx.helper as oh import onnx.numpy_helper as onh +import onnx_ir as ir -from onnxscript import ir from onnxscript.rewriter import pattern diff --git a/onnxscript/_framework_apis/torch_2_5.py b/onnxscript/_framework_apis/torch_2_5.py index 9f31934701..6187429a90 100644 --- a/onnxscript/_framework_apis/torch_2_5.py +++ b/onnxscript/_framework_apis/torch_2_5.py @@ -18,7 +18,9 @@ import pathlib from typing import Callable -from onnxscript import ir, optimizer, version_converter +import onnx_ir as ir + +from onnxscript import optimizer, version_converter from onnxscript.function_libs.torch_lib import registration diff --git a/onnxscript/_framework_apis/torch_2_6.py b/onnxscript/_framework_apis/torch_2_6.py index 2d166cb967..be3b892eb8 100644 --- a/onnxscript/_framework_apis/torch_2_6.py +++ b/onnxscript/_framework_apis/torch_2_6.py @@ -15,7 +15,9 @@ import logging from typing import TYPE_CHECKING -from onnxscript import ir, optimizer, version_converter +import onnx_ir as ir + +from onnxscript import optimizer, version_converter from onnxscript._framework_apis.torch_2_5 import ( check_model, get_torchlib_ops, diff --git a/onnxscript/_internal/autocast.py b/onnxscript/_internal/autocast.py index 1177882abc..38ae53714d 100644 --- a/onnxscript/_internal/autocast.py +++ b/onnxscript/_internal/autocast.py @@ -7,8 +7,9 @@ import numpy as np import onnx +import onnx_ir as ir -from onnxscript import ir, tensor +from onnxscript import tensor if TYPE_CHECKING: from onnxscript._internal import converter diff --git a/onnxscript/_internal/main.py b/onnxscript/_internal/main.py index 804dbfd135..0db97344ad 100644 --- a/onnxscript/_internal/main.py +++ b/onnxscript/_internal/main.py @@ -8,10 +8,10 @@ import sys from typing import Any, Callable, Optional, Sequence, TypeVar +import onnx_ir as ir from typing_extensions import ParamSpec import onnxscript -from onnxscript import ir from onnxscript._internal import ast_utils, converter, irbuilder, values _R = TypeVar("_R") diff --git a/onnxscript/_internal/param_manipulation.py b/onnxscript/_internal/param_manipulation.py index c75d42504b..1bc32decc4 100644 --- a/onnxscript/_internal/param_manipulation.py +++ b/onnxscript/_internal/param_manipulation.py @@ -7,7 +7,7 @@ import collections from typing import Any, OrderedDict -from onnxscript import ir +import onnx_ir as ir def separate_input_attributes_from_arguments( diff --git a/onnxscript/_internal/param_manipulation_test.py b/onnxscript/_internal/param_manipulation_test.py index 892ecd28e7..81b6e83e90 100644 --- a/onnxscript/_internal/param_manipulation_test.py +++ b/onnxscript/_internal/param_manipulation_test.py @@ -5,9 +5,9 @@ import collections import unittest +import onnx_ir as ir import parameterized -from onnxscript import ir from onnxscript._internal import param_manipulation TEST_INPUT = "TEST_INPUT" diff --git a/onnxscript/backend/onnx_backend.py b/onnxscript/backend/onnx_backend.py index ef93bb50b7..9cca75d45b 100644 --- a/onnxscript/backend/onnx_backend.py +++ b/onnxscript/backend/onnx_backend.py @@ -9,6 +9,7 @@ import numpy as np import onnx import onnx.numpy_helper +import onnx_ir as ir from onnx.backend.test import __file__ as backend_folder from onnxscript.backend import onnx_export @@ -99,10 +100,35 @@ def _load(folder, names): res.append(t) return res + @staticmethod + def _normalize_tensor(value): + if isinstance(value, onnx.TensorProto): + return ir.tensor(value).numpy() + return value + def __repr__(self): """Usual""" return f"{self.__class__.__name__}({self.folder!r})" + @classmethod + def from_test_case(cls, test_case): + if test_case.model is None: + raise ValueError(f"Test case {test_case.name!r} does not define a model.") + if test_case.data_sets is None: + raise ValueError(f"Test case {test_case.name!r} does not define test data.") + obj = cls.__new__(cls) + obj.folder = test_case.name + obj.onnx_path = None + obj.onnx_model = test_case.model + obj.tests = [ + dict( + inputs=[cls._normalize_tensor(value) for value in inputs], + outputs=[cls._normalize_tensor(value) for value in outputs], + ) + for inputs, outputs in test_case.data_sets + ] + return obj + def __init__(self, folder): if not os.path.exists(folder): raise FileNotFoundError(f"Unable to find folder {folder!r}.") # pragma: no cover @@ -288,6 +314,14 @@ def enumerate_onnx_tests(series, fct_filter=None) -> Iterator[OnnxBackendTest]: root = os.path.dirname(backend_folder) sub = os.path.join(root, "data", series) if not os.path.exists(sub): + if series == "node": + from onnx.backend.test.loader import load_model_tests + + for test_case in load_model_tests(kind=series): + if fct_filter is not None and not fct_filter(test_case.name): + continue + yield OnnxBackendTest.from_test_case(test_case) + return raise FileNotFoundError( f"Unable to find series of tests in {root!r}, subfolders:\n" + "\n".join(os.listdir(root)) diff --git a/onnxscript/function_libs/tools/torch_lib/deduce_type_constraints.py b/onnxscript/function_libs/tools/torch_lib/deduce_type_constraints.py index 37d358878e..0314c477be 100644 --- a/onnxscript/function_libs/tools/torch_lib/deduce_type_constraints.py +++ b/onnxscript/function_libs/tools/torch_lib/deduce_type_constraints.py @@ -9,9 +9,9 @@ import onnx import onnx.defs +import onnx_ir as ir import onnxscript -from onnxscript import ir logger = logging.getLogger(__name__) diff --git a/onnxscript/function_libs/torch_lib/ops/common.py b/onnxscript/function_libs/torch_lib/ops/common.py index 38544b59ba..4f9d5c00ac 100644 --- a/onnxscript/function_libs/torch_lib/ops/common.py +++ b/onnxscript/function_libs/torch_lib/ops/common.py @@ -9,10 +9,11 @@ import numpy.typing as npt import onnx +import onnx_ir as ir import onnxscript import onnxscript.values -from onnxscript import BOOL, INT64, ir +from onnxscript import BOOL, INT64 from onnxscript import opset18 as op from onnxscript.function_libs.torch_lib import _constants, tensor_typing from onnxscript.function_libs.torch_lib.tensor_typing import RealType diff --git a/onnxscript/function_libs/torch_lib/ops/nn.py b/onnxscript/function_libs/torch_lib/ops/nn.py index 21807c06d5..0b8d6cba50 100644 --- a/onnxscript/function_libs/torch_lib/ops/nn.py +++ b/onnxscript/function_libs/torch_lib/ops/nn.py @@ -17,7 +17,9 @@ import math from typing import Optional, Sequence, Tuple, TypeVar, Union -from onnxscript import BFLOAT16, BOOL, DOUBLE, FLOAT, FLOAT16, INT64, ir +import onnx_ir as ir + +from onnxscript import BFLOAT16, BOOL, DOUBLE, FLOAT, FLOAT16, INT64 from onnxscript.function_libs.torch_lib.registration import torch_op from onnxscript.function_libs.torch_lib.tensor_typing import ( IntType, diff --git a/onnxscript/function_libs/torch_lib/ops/quantized_decomposed.py b/onnxscript/function_libs/torch_lib/ops/quantized_decomposed.py index 7d77208f84..b52dd2b6ef 100644 --- a/onnxscript/function_libs/torch_lib/ops/quantized_decomposed.py +++ b/onnxscript/function_libs/torch_lib/ops/quantized_decomposed.py @@ -13,7 +13,8 @@ from typing import Optional -from onnxscript import ir +import onnx_ir as ir + from onnxscript.function_libs.torch_lib.ops import common from onnxscript.function_libs.torch_lib.registration import torch_op from onnxscript.onnx_opset import opset18 as op diff --git a/onnxscript/ir/_schemas_test.py b/onnxscript/ir/_schemas_test.py index 82082d031f..79383bb5fa 100644 --- a/onnxscript/ir/_schemas_test.py +++ b/onnxscript/ir/_schemas_test.py @@ -5,10 +5,11 @@ import unittest from typing import Any, Optional, Sequence, TypeVar, Union +import onnx_ir as ir import parameterized import onnxscript -from onnxscript import FLOAT, INT64, ir +from onnxscript import FLOAT, INT64 from onnxscript.ir import _schemas _TestTypeVarConstraints = TypeVar("_TestTypeVarConstraints", INT64, FLOAT) diff --git a/onnxscript/optimizer/__init__.py b/onnxscript/optimizer/__init__.py index b8e1d03808..047da2f15a 100644 --- a/onnxscript/optimizer/__init__.py +++ b/onnxscript/optimizer/__init__.py @@ -16,10 +16,10 @@ ] import onnx +import onnx_ir as ir import onnx_ir.passes.common as common_passes import onnxscript.optimizer._constant_folding as constant_folding -from onnxscript import ir from onnxscript.optimizer._constant_folding import FOLDED_FROM_KEY, basic_constant_propagation from onnxscript.optimizer._constant_folding import fold_constants as fold_constants_ir from onnxscript.optimizer._optimizer import optimize_ir diff --git a/onnxscript/optimizer/_constant_folding_test.py b/onnxscript/optimizer/_constant_folding_test.py index e4e92619e2..ad04b1e513 100644 --- a/onnxscript/optimizer/_constant_folding_test.py +++ b/onnxscript/optimizer/_constant_folding_test.py @@ -6,10 +6,10 @@ import numpy as np import onnx +import onnx_ir as ir import parameterized import onnxscript.optimizer as optimizer -from onnxscript import ir from onnxscript.optimizer import _constant_folding diff --git a/onnxscript/optimizer/_function_folding_test.py b/onnxscript/optimizer/_function_folding_test.py index 6f2b052b9e..3314cee621 100644 --- a/onnxscript/optimizer/_function_folding_test.py +++ b/onnxscript/optimizer/_function_folding_test.py @@ -3,9 +3,10 @@ import unittest import onnx +import onnx_ir as ir import onnxscript.testing -from onnxscript import ir, optimizer +from onnxscript import optimizer def _create_model(model_text: str) -> ir.Model: diff --git a/onnxscript/rewriter/__init__.py b/onnxscript/rewriter/__init__.py index 4095e99f05..e70e2c7f1a 100644 --- a/onnxscript/rewriter/__init__.py +++ b/onnxscript/rewriter/__init__.py @@ -22,9 +22,9 @@ ] import onnx +import onnx_ir as ir import onnx_ir.passes.common as common_passes -from onnxscript import ir from onnxscript.rewriter import pattern from onnxscript.rewriter._basics import MatchContext, MatchingTracer, MatchResult, MatchStatus from onnxscript.rewriter._rewrite_rule import ( diff --git a/onnxscript/rewriter/_basics.py b/onnxscript/rewriter/_basics.py index 9b66ff49e6..ad3dff0ef7 100644 --- a/onnxscript/rewriter/_basics.py +++ b/onnxscript/rewriter/_basics.py @@ -9,7 +9,7 @@ from collections import defaultdict from typing import TYPE_CHECKING, Any, MutableSequence, Sequence, Union -from onnxscript import ir +import onnx_ir as ir if TYPE_CHECKING: import onnxscript.rewriter._pattern_ir as _pattern_ir diff --git a/onnxscript/rewriter/_context_test.py b/onnxscript/rewriter/_context_test.py index 4b660dcc3e..fa8f1a4b88 100644 --- a/onnxscript/rewriter/_context_test.py +++ b/onnxscript/rewriter/_context_test.py @@ -6,7 +6,8 @@ import unittest -from onnxscript import ir +import onnx_ir as ir + from onnxscript.rewriter._context import TapeBuilder diff --git a/onnxscript/rewriter/_ir_utils.py b/onnxscript/rewriter/_ir_utils.py index 953d5f33d5..85c796a701 100644 --- a/onnxscript/rewriter/_ir_utils.py +++ b/onnxscript/rewriter/_ir_utils.py @@ -6,8 +6,9 @@ from typing import Callable, Sequence import numpy as np +import onnx_ir as ir -from onnxscript import ir, optimizer +from onnxscript import optimizer def display_nodes(nodes: Sequence[ir.Node]) -> None: diff --git a/onnxscript/rewriter/_matcher.py b/onnxscript/rewriter/_matcher.py index f54b77033f..36f1ba14f4 100644 --- a/onnxscript/rewriter/_matcher.py +++ b/onnxscript/rewriter/_matcher.py @@ -12,9 +12,10 @@ Sequence, ) +import onnx_ir as ir + import onnxscript.rewriter._basics as _basics import onnxscript.rewriter._pattern_ir as _pattern_ir -from onnxscript import ir def _valid_to_replace( diff --git a/onnxscript/rewriter/_pattern_ir.py b/onnxscript/rewriter/_pattern_ir.py index 03749fa0a4..73de2eb817 100644 --- a/onnxscript/rewriter/_pattern_ir.py +++ b/onnxscript/rewriter/_pattern_ir.py @@ -20,8 +20,9 @@ Union, ) +import onnx_ir as ir + import onnxscript.rewriter._basics as _basics -from onnxscript import ir T = TypeVar("T") diff --git a/onnxscript/rewriter/_rewrite_rule.py b/onnxscript/rewriter/_rewrite_rule.py index d1cdf7c5dd..abafb4019f 100644 --- a/onnxscript/rewriter/_rewrite_rule.py +++ b/onnxscript/rewriter/_rewrite_rule.py @@ -13,6 +13,7 @@ TypeVar, ) +import onnx_ir as ir import onnx_ir.passes.common as ir_passes_common import onnxscript.optimizer @@ -22,7 +23,6 @@ import onnxscript.rewriter._matcher as _matcher import onnxscript.rewriter._pattern_ir as _pattern_ir import onnxscript.utils.metadata_merger as metadata_merger -from onnxscript import ir from onnxscript.ir import convenience T = TypeVar("T") diff --git a/onnxscript/rewriter/match_context_test.py b/onnxscript/rewriter/match_context_test.py index e45b8e9ab5..299d74af48 100644 --- a/onnxscript/rewriter/match_context_test.py +++ b/onnxscript/rewriter/match_context_test.py @@ -5,8 +5,8 @@ import unittest import onnx.parser +import onnx_ir as ir -from onnxscript import ir from onnxscript.rewriter import pattern diff --git a/onnxscript/rewriter/onnxruntime/bfloat16_utils/bfloat16_converter.py b/onnxscript/rewriter/onnxruntime/bfloat16_utils/bfloat16_converter.py index 42a3837aa7..34bf993eec 100644 --- a/onnxscript/rewriter/onnxruntime/bfloat16_utils/bfloat16_converter.py +++ b/onnxscript/rewriter/onnxruntime/bfloat16_utils/bfloat16_converter.py @@ -2,7 +2,7 @@ # Licensed under the MIT License. import logging -from onnxscript import ir +import onnx_ir as ir logger = logging.getLogger(__name__) diff --git a/onnxscript/rewriter/onnxruntime/bfloat16_utils/bfloat16_converter_test.py b/onnxscript/rewriter/onnxruntime/bfloat16_utils/bfloat16_converter_test.py index c527855bb7..85456a4906 100644 --- a/onnxscript/rewriter/onnxruntime/bfloat16_utils/bfloat16_converter_test.py +++ b/onnxscript/rewriter/onnxruntime/bfloat16_utils/bfloat16_converter_test.py @@ -5,9 +5,9 @@ import numpy as np import onnx.checker import onnx.shape_inference +import onnx_ir as ir import onnxruntime -from onnxscript import ir from onnxscript.rewriter.onnxruntime.bfloat16_utils import bfloat16_converter diff --git a/onnxscript/rewriter/ort_fusions/fused_matmul_rule_sets.py b/onnxscript/rewriter/ort_fusions/fused_matmul_rule_sets.py index cdc50c99ae..c3de7004ff 100644 --- a/onnxscript/rewriter/ort_fusions/fused_matmul_rule_sets.py +++ b/onnxscript/rewriter/ort_fusions/fused_matmul_rule_sets.py @@ -4,8 +4,9 @@ from typing import ClassVar +import onnx_ir as ir + import onnxscript.rewriter.pattern as orp -from onnxscript import ir from onnxscript.rewriter import _ir_utils diff --git a/onnxscript/rewriter/ort_fusions/fused_matmul_rule_sets_test.py b/onnxscript/rewriter/ort_fusions/fused_matmul_rule_sets_test.py index 43033b9b4f..587d1eae3b 100644 --- a/onnxscript/rewriter/ort_fusions/fused_matmul_rule_sets_test.py +++ b/onnxscript/rewriter/ort_fusions/fused_matmul_rule_sets_test.py @@ -9,11 +9,12 @@ import onnx import onnx.reference import onnx.reference.op_run +import onnx_ir as ir import onnx_ir.passes.common as common_passes import parameterized import onnxscript.rewriter.ort_fusions.fused_matmul_rule_sets as fused_matmul_rule_sets -from onnxscript import FLOAT, ir, script +from onnxscript import FLOAT, script from onnxscript.onnx_opset import opset18 as op from onnxscript.values import Opset diff --git a/onnxscript/rewriter/ort_fusions/group_normalization_merge_silu_test.py b/onnxscript/rewriter/ort_fusions/group_normalization_merge_silu_test.py index dabeaf3851..f9ad4cafd7 100644 --- a/onnxscript/rewriter/ort_fusions/group_normalization_merge_silu_test.py +++ b/onnxscript/rewriter/ort_fusions/group_normalization_merge_silu_test.py @@ -4,8 +4,8 @@ import numpy as np import onnx.parser +import onnx_ir as ir -from onnxscript import ir from onnxscript.rewriter.ort_fusions import ( group_normalization_merge_silu, instance_to_group_normalization, diff --git a/onnxscript/rewriter/ort_fusions/instance_to_group_normalization_test.py b/onnxscript/rewriter/ort_fusions/instance_to_group_normalization_test.py index e5754d78d6..12aef89d5e 100644 --- a/onnxscript/rewriter/ort_fusions/instance_to_group_normalization_test.py +++ b/onnxscript/rewriter/ort_fusions/instance_to_group_normalization_test.py @@ -4,8 +4,8 @@ import numpy as np import onnx.parser +import onnx_ir as ir -from onnxscript import ir from onnxscript.rewriter.ort_fusions import instance_to_group_normalization diff --git a/onnxscript/rewriter/ort_fusions/shape_optimization_test.py b/onnxscript/rewriter/ort_fusions/shape_optimization_test.py index f563ef58d5..b563150e62 100644 --- a/onnxscript/rewriter/ort_fusions/shape_optimization_test.py +++ b/onnxscript/rewriter/ort_fusions/shape_optimization_test.py @@ -4,9 +4,10 @@ import numpy as np import onnx +import onnx_ir as ir import parameterized -from onnxscript import FLOAT, INT64, ir, opset18, script +from onnxscript import FLOAT, INT64, opset18, script from onnxscript.rewriter.ort_fusions import shape_optimization diff --git a/onnxscript/rewriter/ort_fusions/softmax.py b/onnxscript/rewriter/ort_fusions/softmax.py index 10535f57f4..04f58ef9f0 100644 --- a/onnxscript/rewriter/ort_fusions/softmax.py +++ b/onnxscript/rewriter/ort_fusions/softmax.py @@ -5,8 +5,8 @@ import logging import onnx +import onnx_ir as ir -from onnxscript import ir from onnxscript.rewriter._rewrite_rule import RewriteRule, RewriteRuleSet logger = logging.getLogger(__name__) diff --git a/onnxscript/rewriter/ort_fusions/softmax_test.py b/onnxscript/rewriter/ort_fusions/softmax_test.py index e94657d573..59d427e2f1 100644 --- a/onnxscript/rewriter/ort_fusions/softmax_test.py +++ b/onnxscript/rewriter/ort_fusions/softmax_test.py @@ -3,9 +3,9 @@ import unittest import onnx.parser +import onnx_ir as ir import parameterized -from onnxscript import ir from onnxscript.rewriter.ort_fusions import softmax diff --git a/onnxscript/rewriter/pattern_base_test.py b/onnxscript/rewriter/pattern_base_test.py index 77d6521f5e..75b314312d 100644 --- a/onnxscript/rewriter/pattern_base_test.py +++ b/onnxscript/rewriter/pattern_base_test.py @@ -4,7 +4,8 @@ import unittest -from onnxscript import ir +import onnx_ir as ir + from onnxscript.rewriter import pattern from onnxscript.rewriter._basics import MatchFailureError diff --git a/onnxscript/rewriter/pattern_test.py b/onnxscript/rewriter/pattern_test.py index b27e34f0e8..81cb58ff13 100644 --- a/onnxscript/rewriter/pattern_test.py +++ b/onnxscript/rewriter/pattern_test.py @@ -8,10 +8,11 @@ import numpy as np import onnx.checker import onnx.parser +import onnx_ir as ir import onnxscript.optimizer import onnxscript.rewriter -from onnxscript import FLOAT, ir, script +from onnxscript import FLOAT, script from onnxscript import opset17 as op from onnxscript.rewriter import pattern from onnxscript.rewriter.rules.common import _cast_constant_of_shape diff --git a/onnxscript/rewriter/rules/common/_basic_rules.py b/onnxscript/rewriter/rules/common/_basic_rules.py index 7daba40832..1a21f587c9 100644 --- a/onnxscript/rewriter/rules/common/_basic_rules.py +++ b/onnxscript/rewriter/rules/common/_basic_rules.py @@ -12,8 +12,8 @@ from typing import ClassVar, Sequence import numpy as np +import onnx_ir as ir -from onnxscript import ir from onnxscript.rewriter import _ir_utils as ir_utils from onnxscript.rewriter._basics import MatchResult from onnxscript.rewriter._rewrite_rule import RewriteRuleClassBase, RewriteRuleSet diff --git a/onnxscript/rewriter/rules/common/_basic_rules_test.py b/onnxscript/rewriter/rules/common/_basic_rules_test.py index 3eabf1f9b9..9648c32c04 100644 --- a/onnxscript/rewriter/rules/common/_basic_rules_test.py +++ b/onnxscript/rewriter/rules/common/_basic_rules_test.py @@ -8,11 +8,12 @@ import numpy as np import onnx import onnx.reference +import onnx_ir as ir import parameterized import onnxscript import onnxscript.onnx_types as ot -from onnxscript import ir, rewriter +from onnxscript import rewriter from onnxscript.onnx_opset import opset18 from onnxscript.optimizer import _constant_folding, common_passes from onnxscript.rewriter import MatchingTracer, testing diff --git a/onnxscript/rewriter/rules/common/_broadcast_to_matmul.py b/onnxscript/rewriter/rules/common/_broadcast_to_matmul.py index ddf00bc327..ad673a7dcd 100644 --- a/onnxscript/rewriter/rules/common/_broadcast_to_matmul.py +++ b/onnxscript/rewriter/rules/common/_broadcast_to_matmul.py @@ -4,7 +4,8 @@ import logging -from onnxscript import ir +import onnx_ir as ir + from onnxscript.rewriter._rewrite_rule import RewriteRule, RewriteRuleSet logger = logging.getLogger(__name__) diff --git a/onnxscript/rewriter/rules/common/_broadcast_to_matmul_test.py b/onnxscript/rewriter/rules/common/_broadcast_to_matmul_test.py index 4e33544986..8fcafb84a8 100644 --- a/onnxscript/rewriter/rules/common/_broadcast_to_matmul_test.py +++ b/onnxscript/rewriter/rules/common/_broadcast_to_matmul_test.py @@ -6,9 +6,9 @@ import onnx.parser import onnx.shape_inference +import onnx_ir as ir import parameterized -from onnxscript import ir from onnxscript.rewriter.rules.common import _broadcast_to_matmul diff --git a/onnxscript/rewriter/rules/common/_cast_constant_of_shape.py b/onnxscript/rewriter/rules/common/_cast_constant_of_shape.py index 030302f722..8c63d3e777 100644 --- a/onnxscript/rewriter/rules/common/_cast_constant_of_shape.py +++ b/onnxscript/rewriter/rules/common/_cast_constant_of_shape.py @@ -4,7 +4,8 @@ import logging -from onnxscript import ir +import onnx_ir as ir + from onnxscript.rewriter._rewrite_rule import RewriteRule, RewriteRuleSet logger = logging.getLogger(__name__) diff --git a/onnxscript/rewriter/rules/common/_cast_constant_of_shape_test.py b/onnxscript/rewriter/rules/common/_cast_constant_of_shape_test.py index 794491024b..9001f95806 100644 --- a/onnxscript/rewriter/rules/common/_cast_constant_of_shape_test.py +++ b/onnxscript/rewriter/rules/common/_cast_constant_of_shape_test.py @@ -4,8 +4,8 @@ import onnx.checker import onnx.parser +import onnx_ir as ir -from onnxscript import ir from onnxscript.rewriter.rules.common import _cast_constant_of_shape diff --git a/onnxscript/rewriter/rules/common/_collapse_slices.py b/onnxscript/rewriter/rules/common/_collapse_slices.py index 14e7e06d17..712aa410f3 100644 --- a/onnxscript/rewriter/rules/common/_collapse_slices.py +++ b/onnxscript/rewriter/rules/common/_collapse_slices.py @@ -4,7 +4,8 @@ import logging -from onnxscript import ir +import onnx_ir as ir + from onnxscript.rewriter import _ir_utils from onnxscript.rewriter._rewrite_rule import RewriteRule, RewriteRuleSet diff --git a/onnxscript/rewriter/rules/common/_collapse_slices_test.py b/onnxscript/rewriter/rules/common/_collapse_slices_test.py index bcdad4da7e..2c8101a875 100644 --- a/onnxscript/rewriter/rules/common/_collapse_slices_test.py +++ b/onnxscript/rewriter/rules/common/_collapse_slices_test.py @@ -6,8 +6,8 @@ import numpy as np import onnx.parser +import onnx_ir as ir -from onnxscript import ir from onnxscript.rewriter import testing from onnxscript.rewriter.rules.common import _collapse_slices diff --git a/onnxscript/rewriter/rules/common/_fuse_batchnorm.py b/onnxscript/rewriter/rules/common/_fuse_batchnorm.py index 4cd2733463..d8cc8c85ce 100644 --- a/onnxscript/rewriter/rules/common/_fuse_batchnorm.py +++ b/onnxscript/rewriter/rules/common/_fuse_batchnorm.py @@ -18,8 +18,8 @@ from typing import ClassVar, Mapping import numpy as np +import onnx_ir as ir -from onnxscript import ir from onnxscript.rewriter._basics import MatchResult from onnxscript.rewriter._rewrite_rule import RewriteRuleClassBase, RewriteRuleSet diff --git a/onnxscript/rewriter/rules/common/_fuse_batchnorm_test.py b/onnxscript/rewriter/rules/common/_fuse_batchnorm_test.py index a4d1e1efd4..5531080091 100644 --- a/onnxscript/rewriter/rules/common/_fuse_batchnorm_test.py +++ b/onnxscript/rewriter/rules/common/_fuse_batchnorm_test.py @@ -4,9 +4,9 @@ import numpy as np import onnx +import onnx_ir as ir import parameterized -from onnxscript import ir from onnxscript.rewriter import testing from onnxscript.rewriter.rules.common import _fuse_batchnorm diff --git a/onnxscript/rewriter/rules/common/_fuse_conv_affine_test.py b/onnxscript/rewriter/rules/common/_fuse_conv_affine_test.py index d456cab76b..cbe8a8bf41 100644 --- a/onnxscript/rewriter/rules/common/_fuse_conv_affine_test.py +++ b/onnxscript/rewriter/rules/common/_fuse_conv_affine_test.py @@ -3,8 +3,8 @@ import unittest import numpy as np +import onnx_ir as ir -from onnxscript import ir from onnxscript.rewriter import rewrite, testing from onnxscript.rewriter.rules.common import ( affine_conv_fusion_rule, diff --git a/onnxscript/rewriter/rules/common/_gemm_to_matmul_add_test.py b/onnxscript/rewriter/rules/common/_gemm_to_matmul_add_test.py index 90551d8d3b..ae27cac4ef 100644 --- a/onnxscript/rewriter/rules/common/_gemm_to_matmul_add_test.py +++ b/onnxscript/rewriter/rules/common/_gemm_to_matmul_add_test.py @@ -3,8 +3,8 @@ import unittest import onnx.parser +import onnx_ir as ir -from onnxscript import ir from onnxscript.rewriter.rules.common import _gemm_to_matmul_add diff --git a/onnxscript/rewriter/rules/common/_materialize_reshape_shape.py b/onnxscript/rewriter/rules/common/_materialize_reshape_shape.py index 99f3b3b7d6..e7849d7c27 100644 --- a/onnxscript/rewriter/rules/common/_materialize_reshape_shape.py +++ b/onnxscript/rewriter/rules/common/_materialize_reshape_shape.py @@ -14,7 +14,8 @@ from __future__ import annotations -from onnxscript import ir +import onnx_ir as ir + from onnxscript.rewriter import _ir_utils as ir_utils from onnxscript.rewriter._basics import MatchResult from onnxscript.rewriter._rewrite_rule import RewriteRuleClassBase, RewriteRuleSet diff --git a/onnxscript/rewriter/rules/common/_materialize_reshape_shape_test.py b/onnxscript/rewriter/rules/common/_materialize_reshape_shape_test.py index 522d8750c9..d6a4f8a8bc 100644 --- a/onnxscript/rewriter/rules/common/_materialize_reshape_shape_test.py +++ b/onnxscript/rewriter/rules/common/_materialize_reshape_shape_test.py @@ -5,8 +5,8 @@ import unittest import numpy as np +import onnx_ir as ir -from onnxscript import ir from onnxscript.rewriter import testing from onnxscript.rewriter.rules.common import _materialize_reshape_shape diff --git a/onnxscript/rewriter/rules/common/_matmul_add_to_gemm_test.py b/onnxscript/rewriter/rules/common/_matmul_add_to_gemm_test.py index 4c643801fc..ae9bd4f1ce 100644 --- a/onnxscript/rewriter/rules/common/_matmul_add_to_gemm_test.py +++ b/onnxscript/rewriter/rules/common/_matmul_add_to_gemm_test.py @@ -5,10 +5,10 @@ import numpy as np import onnx +import onnx_ir as ir from onnx_ir.passes.common import onnx_checker, shape_inference from parameterized import parameterized -from onnxscript import ir from onnxscript.rewriter import MatchingTracer, MatchStatus, testing from onnxscript.rewriter.rules.common import _matmul_add_to_gemm diff --git a/onnxscript/rewriter/rules/common/_no_op_test.py b/onnxscript/rewriter/rules/common/_no_op_test.py index 2c2f9e6e2b..467cb733e6 100644 --- a/onnxscript/rewriter/rules/common/_no_op_test.py +++ b/onnxscript/rewriter/rules/common/_no_op_test.py @@ -2,9 +2,9 @@ # Licensed under the MIT License. import unittest +import onnx_ir as ir import parameterized -from onnxscript import ir from onnxscript.rewriter.rules.common import _no_op diff --git a/onnxscript/rewriter/rules/common/_remove_expand_before_binary_op.py b/onnxscript/rewriter/rules/common/_remove_expand_before_binary_op.py index bd2501c4b9..cf46fa39b3 100644 --- a/onnxscript/rewriter/rules/common/_remove_expand_before_binary_op.py +++ b/onnxscript/rewriter/rules/common/_remove_expand_before_binary_op.py @@ -13,7 +13,8 @@ from __future__ import annotations -from onnxscript import ir +import onnx_ir as ir + from onnxscript.rewriter._basics import MatchResult from onnxscript.rewriter._ir_utils import get_numpy_value from onnxscript.rewriter._rewrite_rule import RewriteRuleClassBase, RewriteRuleSet diff --git a/onnxscript/rewriter/rules/common/_remove_optional_bias.py b/onnxscript/rewriter/rules/common/_remove_optional_bias.py index db161f0756..0ba2e3426f 100644 --- a/onnxscript/rewriter/rules/common/_remove_optional_bias.py +++ b/onnxscript/rewriter/rules/common/_remove_optional_bias.py @@ -7,8 +7,8 @@ from typing import ClassVar import numpy as np +import onnx_ir as ir -from onnxscript import ir from onnxscript.rewriter._basics import MatchResult from onnxscript.rewriter._rewrite_rule import RewriteRuleClassBase, RewriteRuleSet diff --git a/onnxscript/rewriter/testing.py b/onnxscript/rewriter/testing.py index 2a9d24ee01..2e3d78891f 100644 --- a/onnxscript/rewriter/testing.py +++ b/onnxscript/rewriter/testing.py @@ -7,10 +7,9 @@ import numpy as np import onnx import onnx.reference +import onnx_ir as ir import onnxruntime as ort -from onnxscript import ir - def generate_random_inputs(model: onnx.ModelProto) -> dict[str, Any]: feeds: dict[str, Any] = {} diff --git a/onnxscript/tensor.py b/onnxscript/tensor.py index 6ad8f6bf12..0c70cc57be 100644 --- a/onnxscript/tensor.py +++ b/onnxscript/tensor.py @@ -6,8 +6,9 @@ from typing import Any, Optional import numpy as np +import onnx_ir as ir -from onnxscript import ir, onnx_opset +from onnxscript import onnx_opset from onnxscript._internal import autocast diff --git a/onnxscript/testing/__init__.py b/onnxscript/testing/__init__.py index 0b40a4aa35..2f972cd7ec 100644 --- a/onnxscript/testing/__init__.py +++ b/onnxscript/testing/__init__.py @@ -16,10 +16,10 @@ import google.protobuf.message import numpy as np import onnx +import onnx_ir as ir from onnx import parser import onnxscript -from onnxscript import ir def assert_isomorphic(graph_or_function_1, graph_or_function_2): diff --git a/onnxscript/version_converter/__init__.py b/onnxscript/version_converter/__init__.py index 18cb96b187..ec0ee1ba0c 100644 --- a/onnxscript/version_converter/__init__.py +++ b/onnxscript/version_converter/__init__.py @@ -10,9 +10,9 @@ import logging import onnx +import onnx_ir as ir import onnx_ir.passes.common as common_passes -from onnxscript import ir from onnxscript.version_converter import _c_api_utils, _version_converter logger = logging.getLogger(__name__) diff --git a/onnxscript/version_converter/_c_api_utils.py b/onnxscript/version_converter/_c_api_utils.py index 7f9ac687f4..fd8f4b4f1b 100644 --- a/onnxscript/version_converter/_c_api_utils.py +++ b/onnxscript/version_converter/_c_api_utils.py @@ -7,7 +7,7 @@ import logging from typing import TYPE_CHECKING, Callable, TypeVar -from onnxscript import ir +import onnx_ir as ir if TYPE_CHECKING: import onnx diff --git a/onnxscript/version_converter/_version_converter.py b/onnxscript/version_converter/_version_converter.py index 9b0d941e5f..c08f360e48 100644 --- a/onnxscript/version_converter/_version_converter.py +++ b/onnxscript/version_converter/_version_converter.py @@ -9,11 +9,11 @@ import logging from typing import Callable, Sequence, Union +import onnx_ir as ir import onnx_ir.convenience as ir_convenience import onnx_ir.passes.common as ir_passes_common import onnxscript.utils.metadata_merger as metadata_merger -from onnxscript import ir from onnxscript._internal.tape_builder import BuilderBase, TapeBuilder logger = logging.getLogger(__name__) diff --git a/onnxscript/version_converter/_version_converter_test.py b/onnxscript/version_converter/_version_converter_test.py index 35db893dbf..b12ff96788 100644 --- a/onnxscript/version_converter/_version_converter_test.py +++ b/onnxscript/version_converter/_version_converter_test.py @@ -5,9 +5,10 @@ import unittest import onnx.defs +import onnx_ir as ir import pytest -from onnxscript import ir, version_converter +from onnxscript import version_converter class AdapterCoverageTest(unittest.TestCase): diff --git a/requirements/ci/requirements-onnx-weekly.txt b/requirements/ci/requirements-onnx-weekly.txt index ed5f74b06a..0cc101b7dd 100644 --- a/requirements/ci/requirements-onnx-weekly.txt +++ b/requirements/ci/requirements-onnx-weekly.txt @@ -1 +1 @@ -onnx-weekly==1.22.0.dev20260421 +onnx-weekly==1.23.0.dev20260831 diff --git a/tests/function_libs/torch_lib/ops_test_common.py b/tests/function_libs/torch_lib/ops_test_common.py index 4b5a9c45ec..25d6afbc52 100644 --- a/tests/function_libs/torch_lib/ops_test_common.py +++ b/tests/function_libs/torch_lib/ops_test_common.py @@ -26,6 +26,7 @@ import numpy as np import onnx +import onnx_ir as ir import onnxruntime as ort import onnxruntime.capi.onnxruntime_pybind11_state import pytest @@ -35,7 +36,6 @@ import onnxscript import onnxscript.evaluator -from onnxscript import ir from tests.function_libs.torch_lib import error_reproduction T = TypeVar("T") diff --git a/tests/ir/graph_view_test.py b/tests/ir/graph_view_test.py index 83a51cdaa1..d277d98095 100644 --- a/tests/ir/graph_view_test.py +++ b/tests/ir/graph_view_test.py @@ -4,8 +4,7 @@ import unittest import onnx - -from onnxscript import ir +import onnx_ir as ir class GraphViewTest(unittest.TestCase): diff --git a/tests/ir/serde_roundtrip_test.py b/tests/ir/serde_roundtrip_test.py index 69d23d69e2..5f473177be 100644 --- a/tests/ir/serde_roundtrip_test.py +++ b/tests/ir/serde_roundtrip_test.py @@ -8,10 +8,10 @@ import onnx import onnx.backend.test +import onnx_ir as ir import parameterized import onnxscript.testing -from onnxscript import ir model_folder_path = pathlib.Path(__file__).resolve().parent.parent.parent / "testdata" onnx_backend_test_path = pathlib.Path(onnx.backend.test.__file__).parent / "data" diff --git a/tests/version_converter/version_conversion_test.py b/tests/version_converter/version_conversion_test.py index c012007d12..c3332718f9 100644 --- a/tests/version_converter/version_conversion_test.py +++ b/tests/version_converter/version_conversion_test.py @@ -5,7 +5,9 @@ import pathlib import unittest -from onnxscript import ir, version_converter +import onnx_ir as ir + +from onnxscript import version_converter model_folder_path = pathlib.Path(__file__).resolve().parent.parent.parent / "testdata" diff --git a/tools/ir/model_zoo_test/model_zoo_test.py b/tools/ir/model_zoo_test/model_zoo_test.py index 82d7a54026..2053a517ee 100644 --- a/tools/ir/model_zoo_test/model_zoo_test.py +++ b/tools/ir/model_zoo_test/model_zoo_test.py @@ -18,12 +18,12 @@ import traceback import onnx +import onnx_ir as ir import onnxruntime as ort import tqdm from onnx import hub import onnxscript.testing -from onnxscript import ir def test_model(model_info: hub.ModelInfo) -> float: