已合并
fix: 修改分布式历史遗留用例 #38793
chansinging创建于 6月17日
fix: 修改分布式历史遗留用例 #38793
已合并
共 1 个文件变更+23-23
| @@ -10,6 +10,11 @@ from torch_npu.testing.common_distributed import with_comms, skipIfUnsupportMult | |||
| 10 | 10 | ||
| 11 | class TestRegisterSharding(DTensorTestBase): | 11 | class TestRegisterSharding(DTensorTestBase): |
| 12 | 12 | ||
| 13 | + def trans_BNSD2BSH(self, tensor: torch.Tensor): | ||
| 14 | + tensor = torch.transpose(tensor, 1, 2) | ||
| 15 | + tensor = torch.reshape(tensor, (tensor.shape[0], tensor.shape[1], -1)) | ||
| 16 | + return tensor | ||
| 17 | + | ||
| 13 | def _run_matmul(self, shape1, shape2, device_mesh): | 18 | def _run_matmul(self, shape1, shape2, device_mesh): |
| 14 | x = torch.rand(shape1, device=self.device_type) | 19 | x = torch.rand(shape1, device=self.device_type) |
| 15 | dist_x = distribute_tensor(x, device_mesh, [Replicate()]) | 20 | dist_x = distribute_tensor(x, device_mesh, [Replicate()]) |
| @@ -65,12 +70,12 @@ class TestRegisterSharding(DTensorTestBase): | |||
| 65 | 70 | ||
| 66 | # x is unshardable | 71 | # x is unshardable |
| 67 | local_out, dist_out = self._run_matmul(3, (4, 3, 2), device_mesh) | 72 | local_out, dist_out = self._run_matmul(3, (4, 3, 2), device_mesh) |
| 68 | - self.assertTrue(dist_out.placements[0].is_shard(dim=0)) | 73 | + self.assertTrue(dist_out.placements[0].is_replicate()) |
P | |||
| 69 | self.assertEqual(dist_out.full_tensor(), local_out) | 74 | self.assertEqual(dist_out.full_tensor(), local_out) |
| 70 | 75 | ||
| 71 | # y is unshardable | 76 | # y is unshardable |
| 72 | local_out, dist_out = self._run_matmul((4, 3, 2), 2, device_mesh) | 77 | local_out, dist_out = self._run_matmul((4, 3, 2), 2, device_mesh) |
| 73 | - self.assertTrue(dist_out.placements[0].is_shard(dim=0)) | 78 | + self.assertTrue(dist_out.placements[0].is_replicate()) |
| 74 | self.assertEqual(dist_out.full_tensor(), local_out) | 79 | self.assertEqual(dist_out.full_tensor(), local_out) |
| 75 | 80 | ||
| 76 | 81 | ||
| @@ -87,21 +92,21 @@ class TestRegisterSharding(DTensorTestBase): | |||
| 87 | local_out, dist_out = self._run_matmul( | 92 | local_out, dist_out = self._run_matmul( |
| 88 | 4, (3, 4, 2), device_mesh | 93 | 4, (3, 4, 2), device_mesh |
| 89 | ) # output_shape=(3,2) | 94 | ) # output_shape=(3,2) |
| 90 | - self.assertTrue(dist_out.placements[0].is_partial()) | 95 | + self.assertTrue(dist_out.placements[0].is_replicate()) |
| 91 | self.assertEqual(dist_out.full_tensor(), local_out) | 96 | self.assertEqual(dist_out.full_tensor(), local_out) |
| 92 | 97 | ||
| 93 | # case3: | 98 | # case3: |
| 94 | local_out, dist_out = self._run_matmul( | 99 | local_out, dist_out = self._run_matmul( |
| 95 | 4, (3, 4, 8), device_mesh | 100 | 4, (3, 4, 8), device_mesh |
| 96 | ) # output_shape=(3,8) | 101 | ) # output_shape=(3,8) |
| 97 | - self.assertTrue(dist_out.placements[0].is_shard(dim=-1)) | 102 | + self.assertTrue(dist_out.placements[0].is_replicate()) |
| 98 | self.assertEqual(dist_out.full_tensor(), local_out) | 103 | self.assertEqual(dist_out.full_tensor(), local_out) |
| 99 | 104 | ||
| 100 | # case4: | 105 | # case4: |
| 101 | local_out, dist_out = self._run_matmul( | 106 | local_out, dist_out = self._run_matmul( |
| 102 | 4, (8, 4, 3), device_mesh | 107 | 4, (8, 4, 3), device_mesh |
| 103 | ) # output_shape=(8,3) | 108 | ) # output_shape=(8,3) |
| 104 | - self.assertTrue(dist_out.placements[0].is_shard(dim=0)) | 109 | + self.assertTrue(dist_out.placements[0].is_replicate()) |
| 105 | self.assertEqual(dist_out.full_tensor(), local_out) | 110 | self.assertEqual(dist_out.full_tensor(), local_out) |
| 106 | 111 | ||
| 107 | # ...nm@m=...n | 112 | # ...nm@m=...n |
| @@ -109,14 +114,14 @@ class TestRegisterSharding(DTensorTestBase): | |||
| 109 | local_out, dist_out = self._run_matmul( | 114 | local_out, dist_out = self._run_matmul( |
| 110 | (8, 2, 4), 4, device_mesh | 115 | (8, 2, 4), 4, device_mesh |
| 111 | ) # output_shape=(8,2) | 116 | ) # output_shape=(8,2) |
| 112 | - self.assertTrue(dist_out.placements[0].is_shard(dim=0)) | 117 | + self.assertTrue(dist_out.placements[0].is_replicate()) |
| 113 | self.assertEqual(dist_out.full_tensor(), local_out) | 118 | self.assertEqual(dist_out.full_tensor(), local_out) |
| 114 | 119 | ||
| 115 | # case6: | 120 | # case6: |
| 116 | local_out, dist_out = self._run_matmul( | 121 | local_out, dist_out = self._run_matmul( |
| 117 | (2, 4, 8), 8, device_mesh | 122 | (2, 4, 8), 8, device_mesh |
| 118 | ) # output_shape=(2,4) | 123 | ) # output_shape=(2,4) |
| 119 | - self.assertTrue(dist_out.placements[0].is_partial()) | 124 | + self.assertTrue(dist_out.placements[0].is_replicate()) |
| 120 | self.assertEqual(dist_out.full_tensor(), local_out) | 125 | self.assertEqual(dist_out.full_tensor(), local_out) |
| 121 | 126 | ||
| 122 | # ...nm@...mk=...nk(braodcast) | 127 | # ...nm@...mk=...nk(braodcast) |
| @@ -124,21 +129,21 @@ class TestRegisterSharding(DTensorTestBase): | |||
| 124 | local_out, dist_out = self._run_matmul( | 129 | local_out, dist_out = self._run_matmul( |
| 125 | (2, 8, 4), (2, 2, 4, 2), device_mesh | 130 | (2, 8, 4), (2, 2, 4, 2), device_mesh |
| 126 | ) # output_shape=(2,2,8,2) | 131 | ) # output_shape=(2,2,8,2) |
| 127 | - self.assertTrue(dist_out.placements[0].is_shard(dim=2)) | 132 | + self.assertTrue(dist_out.placements[0].is_replicate()) |
| 128 | self.assertEqual(dist_out.full_tensor(), local_out) | 133 | self.assertEqual(dist_out.full_tensor(), local_out) |
| 129 | 134 | ||
| 130 | # case8: max_size_in_2 and not max_dim1_index==len(shape2)-2: | 135 | # case8: max_size_in_2 and not max_dim1_index==len(shape2)-2: |
| 131 | local_out, dist_out = self._run_matmul( | 136 | local_out, dist_out = self._run_matmul( |
| 132 | (2, 4), (8, 2, 4, 2), device_mesh | 137 | (2, 4), (8, 2, 4, 2), device_mesh |
| 133 | ) # output_shape=(8,2,2,2) | 138 | ) # output_shape=(8,2,2,2) |
| 134 | - self.assertTrue(dist_out.placements[0].is_shard(dim=0)) | 139 | + self.assertTrue(dist_out.placements[0].is_replicate()) |
| 135 | self.assertEqual(dist_out.full_tensor(), local_out) | 140 | self.assertEqual(dist_out.full_tensor(), local_out) |
| 136 | 141 | ||
| 137 | # case9: sharding the core dimension | 142 | # case9: sharding the core dimension |
| 138 | local_out, dist_out = self._run_matmul( | 143 | local_out, dist_out = self._run_matmul( |
| 139 | (2, 2, 4), (2, 2, 4, 2), device_mesh | 144 | (2, 2, 4), (2, 2, 4, 2), device_mesh |
| 140 | ) # output_shape=(2,2,2,2) | 145 | ) # output_shape=(2,2,2,2) |
| 141 | - self.assertTrue(dist_out.placements[0].is_partial()) | 146 | + self.assertTrue(dist_out.placements[0].is_replicate()) |
| 142 | self.assertEqual(dist_out.full_tensor(), local_out) | 147 | self.assertEqual(dist_out.full_tensor(), local_out) |
| 143 | 148 | ||
| 144 | 149 | ||
| @@ -210,23 +215,17 @@ class TestRegisterSharding(DTensorTestBase): | |||
| 210 | value = torch.randn(1, 32, 128, 128, device=self.device_type, dtype=torch.float32) | 215 | value = torch.randn(1, 32, 128, 128, device=self.device_type, dtype=torch.float32) |
| 211 | dy = torch.randn(1, 32, 128, 128, device=self.device_type, dtype=torch.float32) | 216 | dy = torch.randn(1, 32, 128, 128, device=self.device_type, dtype=torch.float32) |
| 212 | 217 | ||
| 213 | - # get attention_in | ||
| 214 | - query = torch.matmul(query, key.transpose(2, 3)).mul(scale) | ||
| 215 | - softmax_res, x_max, x_sum = self.tsoftmax(query.to(torch.float32)) | ||
| 216 | - attention_in = torch.matmul(softmax_res, value) | ||
| 217 | - | ||
| 218 | query = self.trans_BNSD2BSH(query) | 218 | query = self.trans_BNSD2BSH(query) |
| 219 | key = self.trans_BNSD2BSH(key) | 219 | key = self.trans_BNSD2BSH(key) |
| 220 | value = self.trans_BNSD2BSH(value) | 220 | value = self.trans_BNSD2BSH(value) |
| 221 | dy = self.trans_BNSD2BSH(dy) | 221 | dy = self.trans_BNSD2BSH(dy) |
| 222 | 222 | ||
| 223 | - x_max = x_max.expand(1, 32, 128, 8).npu() | 223 | + out, x_max, x_sum, _, _, _, _ = torch_npu.npu_fusion_attention( |
| 224 | - x_sum = x_sum.expand(1, 32, 128, 8).npu() | 224 | + query, key, value, head_num=32, input_layout="BSH", scale=scale) |
| 225 | - out = self.trans_BNSD2BSH(attention_in) | ||
| 226 | 225 | ||
| 227 | - dq, dk, dv, dpse = torch_npu.npu_fusion_attention_grad( | 226 | + dq, dk, dv, dpse, *_ = torch_npu.npu_fusion_attention_grad( |
| 228 | query, key, value, dy, head_num=32, input_layout="BSH", | 227 | query, key, value, dy, head_num=32, input_layout="BSH", |
| 229 | - softmax_max=x_max, softmax_sum=x_sum, attention_in=attention_in, scale_value=scale) | 228 | + softmax_max=x_max, softmax_sum=x_sum, attention_in=out, scale_value=scale) |
| 230 | 229 | ||
| 231 | device_mesh = self.build_device_mesh() | 230 | device_mesh = self.build_device_mesh() |
| 232 | dist_query = distribute_tensor(query, device_mesh, [Replicate()]) | 231 | dist_query = distribute_tensor(query, device_mesh, [Replicate()]) |
| @@ -235,10 +234,11 @@ class TestRegisterSharding(DTensorTestBase): | |||
| 235 | dist_dy = distribute_tensor(dy, device_mesh, [Replicate()]) | 234 | dist_dy = distribute_tensor(dy, device_mesh, [Replicate()]) |
| 236 | dist_xmax = distribute_tensor(x_max, device_mesh, [Replicate()]) | 235 | dist_xmax = distribute_tensor(x_max, device_mesh, [Replicate()]) |
| 237 | dist_xsum = distribute_tensor(x_sum, device_mesh, [Replicate()]) | 236 | dist_xsum = distribute_tensor(x_sum, device_mesh, [Replicate()]) |
| 238 | - dist_attention_in = distribute_tensor(out, device_mesh, [Replicate()]) | 237 | + dist_out = distribute_tensor(out, device_mesh, [Replicate()]) |
| 239 | - dist_dq, dist_dk, dist_dv, dist_dpse = torch_npu.npu_fusion_attention_grad( | 238 | + |
| 239 | + dist_dq, dist_dk, dist_dv, dist_dpse, *_ = torch_npu.npu_fusion_attention_grad( | ||
| 240 | dist_query, dist_key, dist_value, dist_dy, head_num=32, input_layout="BSH", | 240 | dist_query, dist_key, dist_value, dist_dy, head_num=32, input_layout="BSH", |
| 241 | - softmax_max=dist_xmax, softmax_sum=dist_xsum, attention_in=dist_attention_in, scale_value=scale) | 241 | + softmax_max=dist_xmax, softmax_sum=dist_xsum, attention_in=dist_out, scale_value=scale) |
| 242 | 242 | ||
| 243 | self.assertEqual(dist_dq.full_tensor(), dq) | 243 | self.assertEqual(dist_dq.full_tensor(), dq) |
| 244 | self.assertEqual(dist_dk.full_tensor(), dk) | 244 | self.assertEqual(dist_dk.full_tensor(), dk) |
这些断言逻辑调整的原因是什么,之前都没跑过吗