已合并
fix_device_dynamo_trace #31558
Ambi创建于 3月9日
fix_device_dynamo_trace #31558
已合并
Ambi创建于 3月9日
已删除 :v2.8.0合入到Ascend/pytorchv2.8.0
3 个文件变更+53-1
Atest/_inductor/test_compile_autograd.py+52-0
@@ -0,0 +1,52 @@
1+import os
2+import torch
3+from torch.testing._internal.common_utils import run_tests, parametrize, instantiate_parametrized_tests
4+from testutils import TestUtils
5+import torch_npu
6+import torch_npu._inductor
7+ 
8+ 
9+class DeviceCheckFunc(torch.autograd.Function):
10+ @staticmethod
11+ def forward(ctx, x):
12+ return x.clone()
13+ 
14+ @staticmethod
15+ def backward(ctx, grad_output):
16+ # 在 backward 中访问 device 对象
17+ dev = torch.device("npu")
18+ return grad_output
19+ 
20+ 
21+class TestCompiledAutograd(TestUtils):
22+ 
23+ @parametrize("input_dim", [10])
24+ @parametrize("device", ["npu"])
25+ def test_compiled_autograd(self, input_dim, device):
26+ 
27+ def model_fn(x):
28+ return DeviceCheckFunc.apply(x).sum()
29+ 
30+ torch.manual_seed(42)
31+ 
32+ x = torch.randn(input_dim, requires_grad=True, device=device)
33+ 
34+ with torch._dynamo.utils.maybe_enable_compiled_autograd(True):
35+ compiled_fn = torch.compile(
36+ model_fn,
37+ backend="inductor",
38+ dynamic=False
39+ )
40+ 
41+ loss = compiled_fn(x)
42+ loss.backward()
43+ 
44+ grad = x.grad.detach().cpu().numpy()
45+ print("grad:", grad)
46+ 
47+ 
48+instantiate_parametrized_tests(TestCompiledAutograd)
49+ 
50+ 
51+if __name__ == "__main__":
52+ run_tests()
Mtest/_inductor/test_dvm_graph_fusion.py+1-0
@@ -26,6 +26,7 @@ class MatMulModule(torch.nn.Module):
26 26 
27 def forward(self, a, b):27 def forward(self, a, b):
28 mm = torch.mm(a.t(), b)28 mm = torch.mm(a.t(), b)
29+ mm = mm.to(torch.float32)
29 return mm + 330 return mm + 3
30 31 
31 32 
Mtorch_npu/utils/_dynamo.py+0-1
@@ -74,7 +74,6 @@ def UserDefinedClassVariable__new__(cls, value, **kwargs):
74 torch_npu.npu.LongTensor,74 torch_npu.npu.LongTensor,
75 torch_npu.npu.ShortTensor,75 torch_npu.npu.ShortTensor,
76 torch_npu.npu.BFloat16Tensor,76 torch_npu.npu.BFloat16Tensor,
77- torch.device,
78 ]:77 ]:
79 return TorchInGraphFunctionVariable(value, **kwargs)78 return TorchInGraphFunctionVariable(value, **kwargs)
80 return cls.__new__raw(cls)79 return cls.__new__raw(cls)