已合并
add test for nn.Module.npu() #35163
add test for nn.Module.npu() #35163
已合并
zf_zhang创建于 5月9日
1 个文件变更+93-68
Mtest/nn/test_nn_api.py+93-68
@@ -1,68 +1,93 @@
1-"""1+# Owner(s): ["module: nn"]
2-Add validation cases for torch.nn APIs on NPU:2+ 
3-1. test/test_nn.py from PyTorch community lacks sufficient API validations, so this file is added.3+"""
4-2. This file validates torch.nn.Parameter, torch.nn.Buffer (extendable).4+Add validation cases for torch.nn APIs on NPU:
5-"""5+1. test/test_nn.py from PyTorch community lacks sufficient API validations, so this file is added.
6- 6+2. This file validates torch.nn.Parameter, torch.nn.Buffer, torch.nn.Module.npu (extendable).
7-import torch7+"""
8-import torch.nn as nn8+ 
9-from torch.testing._internal.common_utils import run_tests, TestCase9+import torch
10-import torch_npu10+import torch.nn as nn
11- 11+from torch.testing._internal.common_utils import run_tests, TestCase
12-device = torch.device("npu:0")12+ 
13- 13+ 
14- 14+device = torch.device("npu:0")
15-class TestNPUParameterBuffer(TestCase):15+ 
16- def test_parameter_api(self):16+ 
17- """验证torch.nn.Parameter的创建、属性及NPU设备支持"""17+class TestNPUParameterBuffer(TestCase):
18- # 直接在NPU上创建Parameter(推荐方式)18+ def test_parameter_api(self):
19- p = nn.Parameter(torch.randn(10, 20, device=device))19+ """Verifies Parameter creation, attributes, and in-place modification on NPU."""
20- 20+ # 直接在NPU上创建Parameter(推荐方式)
21- # 验证类型和属性21+ p = nn.Parameter(torch.randn(10, 20, device=device))
22- self.assertIsInstance(p, nn.Parameter)22+ 
23- self.assertTrue(p.requires_grad)23+ # 验证类型和属性
24- self.assertEqual(p.device, device)24+ self.assertIsInstance(p, nn.Parameter)
25- self.assertEqual(p.shape, (10, 20))25+ self.assertTrue(p.requires_grad)
26- 26+ self.assertEqual(p.device, device)
27- # 验证数据操作27+ self.assertEqual(p.shape, (10, 20))
28- p.data = p.data * 2 # Parameter支持.data访问28+ 
29- self.assertEqual(p.shape, (10, 20)) # 形状不变29+ # 验证数据操作
30- 30+ p.data = p.data * 2 # Parameter支持.data访问
31- # 修改requires_grad31+ self.assertEqual(p.shape, (10, 20)) # 形状不变
32- p.requires_grad = False32+ 
33- self.assertFalse(p.requires_grad)33+ # 修改requires_grad
34- 34+ p.requires_grad = False
35- def test_buffer_api(self):35+ self.assertFalse(p.requires_grad)
36- """验证torch.nn.Buffer的创建、persistent属性及NPU设备支持"""36+ 
37- # 创建persistent Buffer(默认)37+ def test_buffer_api(self):
38- b1 = nn.Buffer(torch.randn(5, 5, device=device))38+ """Verifies Buffer creation, attributes, and in-place modification on NPU."""
39- self.assertIsInstance(b1, nn.Buffer)39+ # 创建persistent Buffer(默认)
40- self.assertFalse(b1.requires_grad) # Buffer默认不需要梯度40+ b1 = nn.Buffer(torch.randn(5, 5, device=device))
41- self.assertEqual(b1.device, device)41+ self.assertIsInstance(b1, nn.Buffer)
42- 42+ self.assertFalse(b1.requires_grad) # Buffer默认不需要梯度
43- # 创建non-persistent Buffer43+ self.assertEqual(b1.device, device)
44- b2 = nn.Buffer(torch.randn(3, 3, device=device), persistent=False)44+ 
45- self.assertFalse(b2.requires_grad)45+ # 创建non-persistent Buffer
46- self.assertEqual(b2.device, device)46+ b2 = nn.Buffer(torch.randn(3, 3, device=device), persistent=False)
47- 47+ self.assertFalse(b2.requires_grad)
48- # 验证persistent属性在Module中的行为48+ self.assertEqual(b2.device, device)
49- class TestModule(nn.Module):49+ 
50- def __init__(self):50+ # 验证persistent属性在Module中的行为
51- super().__init__()51+ class TestModule(nn.Module):
52- self.register_buffer('persistent_buf', b1)52+ def __init__(self):
53- self.register_buffer('non_persistent_buf', b2, persistent=False)53+ super().__init__()
54- 54+ self.register_buffer("persistent_buf", b1)
55- m = TestModule()55+ self.register_buffer("non_persistent_buf", b2, persistent=False)
56- state_dict = m.state_dict()56+ 
57- 57+ m = TestModule()
58- # persistent buffer应在state_dict58+ state_dict = m.state_dict()
59- self.assertIn('persistent_buf', state_dict)59+ 
60- # non-persistent buffer应在state_dict中 60+ # persistent buffer应在state_dict中
61- self.assertNotIn('non_persistent_buf', state_dict)61+ self.assertIn("persistent_buf", state_dict)
62- 62+ # non-persistent buffer不应在state_dict中
63- # 验证Buffer在state_dict中的值正确63+ self.assertNotIn("non_persistent_buf", state_dict)
64- self.assertEqual(state_dict['persistent_buf'].device, device)64+ 
65- 65+ # 验证Buffer在state_dict中的值正确
66- 66+ self.assertEqual(state_dict["persistent_buf"].device, device)
67-if __name__ == "__main__":67+ 
68- run_tests()68+ 
69+class TestNNModuleAPIs(TestCase):
70+ def test_npu(self):
71+ """Verifies that Module.npu() correctly move parameters and buffers to NPU."""
72+ 
73+ class MyModule(nn.Module):
74+ def __init__(self, in_features, out_features):
75+ super().__init__()
76+ self.weight = nn.Parameter(torch.randn(in_features, out_features))
77+ self.register_buffer("buf", torch.randn(out_features))
78+ 
79+ def forward(self, x):
80+ return x @ self.weight + self.buf
81+ 
82+ m = MyModule(3, 5)
83+ self.assertEqual(m.to("npu"), m.npu())
84+ m1 = m.npu()
85+ 
86+ for param in m1.parameters():
87+ self.assertEqual(param.device.type, "npu")
88+ for param in m1.buffers():
89+ self.assertEqual(param.device.type, "npu")
90+ 
91+ 
92+if __name__ == "__main__":
93+ run_tests()