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
3 changes: 2 additions & 1 deletion docs/tutorial/rewriter/examples/broadcast_matmul.py
Original file line number Diff line number Diff line change
Expand Up @@ -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__)
Expand Down
3 changes: 2 additions & 1 deletion docs/tutorial/rewriter/examples/erfgelu.py
Original file line number Diff line number Diff line change
Expand Up @@ -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


Expand Down
2 changes: 1 addition & 1 deletion examples/pattern_matching_example.py
Original file line number Diff line number Diff line change
Expand Up @@ -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


Expand Down
2 changes: 1 addition & 1 deletion examples/pattern_rewriting.py
Original file line number Diff line number Diff line change
Expand Up @@ -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


Expand Down
4 changes: 3 additions & 1 deletion onnxscript/_framework_apis/torch_2_5.py
Original file line number Diff line number Diff line change
Expand Up @@ -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


Expand Down
4 changes: 3 additions & 1 deletion onnxscript/_framework_apis/torch_2_6.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down
3 changes: 2 additions & 1 deletion onnxscript/_internal/autocast.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
2 changes: 1 addition & 1 deletion onnxscript/_internal/main.py
Original file line number Diff line number Diff line change
Expand Up @@ -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")
Expand Down
2 changes: 1 addition & 1 deletion onnxscript/_internal/param_manipulation.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down
2 changes: 1 addition & 1 deletion onnxscript/_internal/param_manipulation_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -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"
Expand Down
34 changes: 34 additions & 0 deletions onnxscript/backend/onnx_backend.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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))
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -9,9 +9,9 @@

import onnx
import onnx.defs
import onnx_ir as ir

import onnxscript
from onnxscript import ir

logger = logging.getLogger(__name__)

Expand Down
3 changes: 2 additions & 1 deletion onnxscript/function_libs/torch_lib/ops/common.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
4 changes: 3 additions & 1 deletion onnxscript/function_libs/torch_lib/ops/nn.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
3 changes: 2 additions & 1 deletion onnxscript/ir/_schemas_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
2 changes: 1 addition & 1 deletion onnxscript/optimizer/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
2 changes: 1 addition & 1 deletion onnxscript/optimizer/_constant_folding_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -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


Expand Down
3 changes: 2 additions & 1 deletion onnxscript/optimizer/_function_folding_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down
2 changes: 1 addition & 1 deletion onnxscript/rewriter/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -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 (
Expand Down
2 changes: 1 addition & 1 deletion onnxscript/rewriter/_basics.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
3 changes: 2 additions & 1 deletion onnxscript/rewriter/_context_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,7 +6,8 @@

import unittest

from onnxscript import ir
import onnx_ir as ir

from onnxscript.rewriter._context import TapeBuilder


Expand Down
3 changes: 2 additions & 1 deletion onnxscript/rewriter/_ir_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down
3 changes: 2 additions & 1 deletion onnxscript/rewriter/_matcher.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down
3 changes: 2 additions & 1 deletion onnxscript/rewriter/_pattern_ir.py
Original file line number Diff line number Diff line change
Expand Up @@ -20,8 +20,9 @@
Union,
)

import onnx_ir as ir

import onnxscript.rewriter._basics as _basics
from onnxscript import ir

T = TypeVar("T")

Expand Down
2 changes: 1 addition & 1 deletion onnxscript/rewriter/_rewrite_rule.py
Original file line number Diff line number Diff line change
Expand Up @@ -13,6 +13,7 @@
TypeVar,
)

import onnx_ir as ir
import onnx_ir.passes.common as ir_passes_common

import onnxscript.optimizer
Expand All @@ -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")
Expand Down
2 changes: 1 addition & 1 deletion onnxscript/rewriter/match_context_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,8 +5,8 @@
import unittest

import onnx.parser
import onnx_ir as ir

from onnxscript import ir
from onnxscript.rewriter import pattern


Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -2,7 +2,7 @@
# Licensed under the MIT License.
import logging

from onnxscript import ir
import onnx_ir as ir

logger = logging.getLogger(__name__)

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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


Expand Down
3 changes: 2 additions & 1 deletion onnxscript/rewriter/ort_fusions/fused_matmul_rule_sets.py
Original file line number Diff line number Diff line change
Expand Up @@ -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


Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down
Loading
Loading