已合并
add test for nn.Module.npuv212 #35264
zf_zhang创建于 5月11日
add test for nn.Module.npuv212 #35264
已合并
共 1 个文件变更+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 torch | 7 | +""" |
| 8 | -import torch.nn as nn | 8 | + |
| 9 | -from torch.testing._internal.common_utils import run_tests, TestCase | 9 | +import torch |
| 10 | -import torch_npu | 10 | +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_grad | 31 | + self.assertEqual(p.shape, (10, 20)) # 形状不变 |
| 32 | - p.requires_grad = False | 32 | + |
| 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 Buffer | 43 | + 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_dict中 | 58 | + 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() | ||