已合并
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
已合并
从已删除 :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() | ||