mirror of
https://github.com/zebrajr/pytorch.git
synced 2025-12-07 12:21:27 +01:00
This reverts commit 1a55fb0ee8.
Reverted https://github.com/pytorch/pytorch/pull/154725 on behalf of https://github.com/malfet due to This added 2nd copy of raise_on_run to common_utils.py which caused lint failures, see https://github.com/pytorch/pytorch/actions/runs/15445374980/job/43473457466 ([comment](https://github.com/pytorch/pytorch/pull/154725#issuecomment-2940503905))
45 lines
1.3 KiB
Python
45 lines
1.3 KiB
Python
# Owner(s): ["oncall: jit"]
|
|
|
|
import torch
|
|
from torch.testing import FileCheck
|
|
from torch.testing._internal.jit_utils import JitTestCase
|
|
|
|
|
|
if __name__ == "__main__":
|
|
raise RuntimeError(
|
|
"This test file is not meant to be run directly, use:\n\n"
|
|
"\tpython test/test_jit.py TESTNAME\n\n"
|
|
"instead."
|
|
)
|
|
|
|
|
|
class TestOpDecompositions(JitTestCase):
|
|
def test_op_decomposition(self):
|
|
def foo(x):
|
|
return torch.var(x, unbiased=True)
|
|
|
|
# TODO: more robust testing
|
|
foo_s = torch.jit.script(foo)
|
|
FileCheck().check("aten::var").run(foo_s.graph)
|
|
torch._C._jit_pass_run_decompositions(foo_s.graph)
|
|
inp = torch.rand([10, 10])
|
|
self.assertEqual(foo(inp), foo_s(inp))
|
|
FileCheck().check_not("aten::var").run(foo_s.graph)
|
|
|
|
def test_registered_decomposition(self):
|
|
@torch.jit.script
|
|
def foo(x):
|
|
return torch.square(x)
|
|
|
|
@torch.jit.script
|
|
def square_decomp(x):
|
|
return torch.pow(x, 2)
|
|
|
|
torch.jit._register_decomposition(
|
|
torch.ops.aten.square.default, square_decomp.graph
|
|
)
|
|
torch._C._jit_pass_run_decompositions(foo.graph)
|
|
FileCheck().check_not("aten::square").check("aten::pow").run(foo.graph)
|
|
x = torch.rand([4])
|
|
self.assertEqual(foo(x), torch.square(x))
|