已合并
add_guard #31897
Ambi创建于 3月17日
add_guard #31897
已合并
Ambi创建于 3月17日
已删除 :v2.8.0合入到Ascend/pytorchv2.8.0
2 个文件变更+35-0
@@ -0,0 +1,33 @@
1+import torch
2+from torch.testing._internal.common_utils import (
3+ run_tests,
4+ instantiate_parametrized_tests,
5+)
6+from testutils import TestUtils
7+import torch_npu
8+ 
9+ 
10+class TestCurrentDevice(TestUtils):
11+ 
12+ def test_npu_current_device(self):
13+ def fn(x):
14+ y = torch.empty(
15+ (2, 3), dtype=torch.float32, device=torch.npu.current_device()
16+ )
17+ y.copy_(x)
18+ return torch.sin(y + y.device.index)
19+ 
20+ counter = torch._dynamo.testing.CompileCounter()
21+ opt_fn = torch.compile(backend=counter, fullgraph=True)(fn)
22+ 
23+ with torch.npu.device(0):
24+ x = torch.randn(2, 3).to("npu")
25+ self.assertEqual(opt_fn(x), fn(x))
26+ self.assertEqual(counter.frame_count, 1)
27+ 
28+ 
29+instantiate_parametrized_tests(TestCurrentDevice)
30+ 
31+ 
32+if __name__ == "__main__":
33+ run_tests()
@@ -72,6 +72,7 @@ torch_c_binding_in_graph_functions_npu = dict.fromkeys(
72 "torch_npu._C._npu_setMemoryFraction",72 "torch_npu._C._npu_setMemoryFraction",
73 "torch_npu._C._npu_synchronize",73 "torch_npu._C._npu_synchronize",
74 "torch_npu._C._npu_resetAccumulatedMemoryStats",74 "torch_npu._C._npu_resetAccumulatedMemoryStats",
75+ "torch_npu._C._npu_setStream",
75 ],76 ],
76 TorchInGraphFunctionVariable,77 TorchInGraphFunctionVariable,
77)78)
@@ -91,6 +92,7 @@ def _patch_npu_trace_rules():
91 torch._dynamo.trace_rules.torch_name_rule_map.append(skip_functions_npu)92 torch._dynamo.trace_rules.torch_name_rule_map.append(skip_functions_npu)
92 torch_module.constant_fold_functions[torch.npu.current_device] = True93 torch_module.constant_fold_functions[torch.npu.current_device] = True
93 torch_module.constant_fold_functions[torch.npu.get_device_properties] = True94 torch_module.constant_fold_functions[torch.npu.get_device_properties] = True
95+ torch_module.constant_fold_functions_need_guards[torch.npu.current_device] = True
94 torch_module.constant_fold_functions[torch.npu.is_available] = True96 torch_module.constant_fold_functions[torch.npu.is_available] = True
95 common_constant_types.add(torch_npu._C._NPUDeviceProperties)97 common_constant_types.add(torch_npu._C._NPUDeviceProperties)
96 98