已关闭
test: adapt ShapeEnv proxy tensor upstream tests #36396
小辉懂编程创建于 5月21日关闭于 5月22日
test: adapt ShapeEnv proxy tensor upstream tests #36396
已关闭
小辉懂编程创建于 5月21日关闭于 5月22日
共 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+@@ -3,6 +3,7 @@
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+@@ -966,7 +967,7 @@ 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+@@ -1831,7 +1832,7 @@ 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
OopenLiBingCI5月21日

此条代码评论区间+27至+29

【openlibing.ci】识别到代码检查告警抑制注释,匹配工具:flake8,请Committer检视其合理性。

likedislike