已合并
test(nn): add test for nn api: torch.nn.Parameter, torch.nn.Buffer #31892
dinglaiping创建于 3月17日
test(nn): add test for nn api: torch.nn.Parameter, torch.nn.Buffer #31892
已合并
dinglaiping创建于 3月17日
已删除 :addtest-nn-api-2.8.0合入到Ascend/pytorchv2.8.0
1 个文件变更+68-0
@@ -0,0 +1,68 @@
1+"""
2+Add validation cases for torch.nn APIs on NPU:
3+1. test/test_nn.py from PyTorch community lacks sufficient API validations, so this file is added.
4+2. This file validates torch.nn.Parameter, torch.nn.Buffer (extendable).
5+"""
6+ 
7+import torch
8+import torch.nn as nn
9+from torch.testing._internal.common_utils import run_tests, TestCase
10+import torch_npu
11+ 
12+device = torch.device("npu:0")
13+ 
14+ 
15+class TestNPUParameterBuffer(TestCase):
16+ def test_parameter_api(self):
17+ """验证torch.nn.Parameter的创建、属性及NPU设备支持"""
18+ # 直接在NPU上创建Parameter(推荐方式)
19+ p = nn.Parameter(torch.randn(10, 20, device=device))
20+
21+ # 验证类型和属性
22+ self.assertIsInstance(p, nn.Parameter)
23+ self.assertTrue(p.requires_grad)
24+ self.assertEqual(p.device, device)
25+ self.assertEqual(p.shape, (10, 20))
26+
27+ # 验证数据操作
28+ p.data = p.data * 2 # Parameter支持.data访问
29+ self.assertEqual(p.shape, (10, 20)) # 形状不变
30+
31+ # 修改requires_grad
32+ p.requires_grad = False
33+ self.assertFalse(p.requires_grad)
34+ 
35+ def test_buffer_api(self):
36+ """验证torch.nn.Buffer的创建、persistent属性及NPU设备支持"""
37+ # 创建persistent Buffer(默认)
38+ b1 = nn.Buffer(torch.randn(5, 5, device=device))
39+ self.assertIsInstance(b1, nn.Buffer)
40+ self.assertFalse(b1.requires_grad) # Buffer默认不需要梯度
41+ self.assertEqual(b1.device, device)
42+
43+ # 创建non-persistent Buffer
44+ b2 = nn.Buffer(torch.randn(3, 3, device=device), persistent=False)
45+ self.assertFalse(b2.requires_grad)
46+ self.assertEqual(b2.device, device)
47+
48+ # 验证persistent属性在Module中的行为
49+ class TestModule(nn.Module):
50+ def __init__(self):
51+ super().__init__()
52+ self.register_buffer('persistent_buf', b1)
53+ self.register_buffer('non_persistent_buf', b2, persistent=False)
54+
55+ m = TestModule()
56+ state_dict = m.state_dict()
57+
58+ # persistent buffer应在state_dict中
59+ self.assertIn('persistent_buf', state_dict)
60+ # non-persistent buffer不应在state_dict中
61+ self.assertNotIn('non_persistent_buf', state_dict)
62+
63+ # 验证Buffer在state_dict中的值正确
64+ self.assertEqual(state_dict['persistent_buf'].device, device)
65+ 
66+ 
67+if __name__ == "__main__":
68+ run_tests()