已合并
fix: test_mapping.py #2167
aijgnem1创建于 2025年4月8日
fix: test_mapping.py #2167
已合并
从refs/pull/2167/head合入到core_r0.8.0
共 1 个文件变更+12-0
| @@ -128,6 +128,10 @@ class TestUnevenSplitGather(DistributedTest): | |||
| 128 | (0, 1), (0, 2), (0, 3)]) | 128 | (0, 1), (0, 2), (0, 3)]) |
| 129 | 129 | ||
| 130 | def test_all_to_all(self, gather_scatter_idx, dtype): | 130 | def test_all_to_all(self, gather_scatter_idx, dtype): |
| 131 | + args = parse_args(None, True) | ||
| 132 | + set_args(args) | ||
| 133 | + destroy_model_parallel() | ||
| 134 | + initialize_model_parallel(tensor_model_parallel_size=self.world_size) | ||
| 131 | input_shapes = [[32, 64, 32], [32, 64, 32, 32], | 135 | input_shapes = [[32, 64, 32], [32, 64, 32, 32], |
| 132 | [29, 51, 57], [29, 51, 65, 57]] | 136 | [29, 51, 57], [29, 51, 65, 57]] |
| 133 | group = mpu.get_tensor_model_parallel_group() | 137 | group = mpu.get_tensor_model_parallel_group() |
| @@ -151,6 +155,10 @@ class TestUnevenSplitGather(DistributedTest): | |||
| 151 | (0, 1), (0, 2), (0, 3)]) | 155 | (0, 1), (0, 2), (0, 3)]) |
| 152 | 156 | ||
| 153 | def test_all_to_all_full_unaligned(self, gather_scatter_idx, dtype): | 157 | def test_all_to_all_full_unaligned(self, gather_scatter_idx, dtype): |
| 158 | + args = parse_args(None, True) | ||
| 159 | + set_args(args) | ||
| 160 | + destroy_model_parallel() | ||
| 161 | + initialize_model_parallel(tensor_model_parallel_size=self.world_size) | ||
| 154 | group = mpu.get_tensor_model_parallel_group() | 162 | group = mpu.get_tensor_model_parallel_group() |
| 155 | rank = dist.get_rank(group) | 163 | rank = dist.get_rank(group) |
| 156 | unscatter_gathered_shapes = [[9, 21, 15], [9, 21, 15, 10]] | 164 | unscatter_gathered_shapes = [[9, 21, 15], [9, 21, 15, 10]] |
| @@ -188,6 +196,10 @@ class TestUnevenSplitGather(DistributedTest): | |||
| 188 | 196 | ||
| 189 | 197 | ||
| 190 | def test_split_gather_default(self, input_shape, dim, dtype): | 198 | def test_split_gather_default(self, input_shape, dim, dtype): |
| 199 | + args = parse_args(None, True) | ||
| 200 | + set_args(args) | ||
| 201 | + destroy_model_parallel() | ||
| 202 | + initialize_model_parallel(tensor_model_parallel_size=self.world_size) | ||
| 191 | group = mpu.get_tensor_model_parallel_group() | 203 | group = mpu.get_tensor_model_parallel_group() |
| 192 | input_tensor = torch.randn(input_shape).cuda().to(dtype) | 204 | input_tensor = torch.randn(input_shape).cuda().to(dtype) |
| 193 | if dim >= input_tensor.dim(): | 205 | if dim >= input_tensor.dim(): |