已关闭
test: adapt ShapeEnv proxy tensor upstream tests #36396
小辉懂编程创建于 5月21日关闭于 5月22日
test: adapt ShapeEnv proxy tensor upstream tests #36396
已关闭
共 1 个文件变更+29-0
| @@ -0,0 +1,29 @@ | |||
| 1 | +diff --git a/test/test_proxy_tensor.py b/test/test_proxy_tensor.py | ||
| 2 | +--- a/test/test_proxy_tensor.py | ||
| 3 | ++++ b/test/test_proxy_tensor.py | ||
| 4 | + | ||
| 5 | + | ||
| 6 | + from torch.testing._internal.common_utils import TestCase, run_tests | ||
| 7 | + import torch | ||
| 8 | ++import torch_npu | ||
| 9 | + import torch._dynamo | ||
| 10 | + import unittest | ||
| 11 | + import warnings | ||
| 12 | + def _get_free_symbols(shape_env): | ||
| 13 | + return len([var for var in vars if var not in shape_env.replacements]) | ||
| 14 | + | ||
| 15 | + def _trace(f, *args): | ||
| 16 | +- inps = [torch.randn(arg) for arg in args] | ||
| 17 | ++ inps = [torch.randn(arg).npu() for arg in args] | ||
| 18 | + return make_fx(f, tracing_mode="symbolic")(*inps) | ||
| 19 | + | ||
| 20 | + # TODO: Need to test the guards themselves specifically as well | ||
| 21 | + class TestSymbolicTracing(TestCase): | ||
| 22 | + def f(a, b): | ||
| 23 | + assert a.shape[0] == b.shape[0] * 2 | ||
| 24 | + return a.cos() | ||
| 25 | +- fx_g = make_fx(f, tracing_mode="symbolic")(torch.randn(16), torch.randn(8)) | ||
| 26 | ++ fx_g = make_fx(f, tracing_mode="symbolic")(torch.randn(16).npu(), torch.randn(8).npu()) | ||
| 27 | + from torch._dynamo.source import LocalSource | ||
| 28 | + self.assertExpectedInline( | ||
| 29 | + str(fx_g.shape_env.produce_guards(fx_placeholder_vals(fx_g), [LocalSource("a"), LocalSource("b")], ignore_static=False)), # noqa: B950 | ||
O | |||
此条代码评论区间+27至+29
【openlibing.ci】识别到代码检查告警抑制注释,匹配工具:flake8,请Committer检视其合理性。