已合并
fix: 修改分布式历史遗留用例 #38793
chansinging创建于 6月17日
fix: 修改分布式历史遗留用例 #38793
已合并
chansinging创建于 6月17日
1 个文件变更+23-23
Mtest/distributed/test_register_sharding.py+23-23
@@ -10,6 +10,11 @@ from torch_npu.testing.common_distributed import with_comms, skipIfUnsupportMult
10 10 
11class TestRegisterSharding(DTensorTestBase):11class 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 unshardable71 # 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
Ppengjingyou6月22日

这些断言逻辑调整的原因是什么,之前都没跑过吗

likedislike
chansinging
chansinging
6月22日 评论:
69 self.assertEqual(dist_out.full_tensor(), local_out)74 self.assertEqual(dist_out.full_tensor(), local_out)
70 75 
71 # y is unshardable76 # 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 @with_comms81 @with_comms
@@ -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_mesh93 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_mesh100 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_mesh107 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=...n112 # ...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_mesh115 (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_mesh122 (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_mesh130 (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_mesh137 (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 dimension142 # 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_mesh144 (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 @with_comms149 @with_comms
@@ -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)