已合并
fix: test_mapping.py #2167
aijgnem1创建于 2025年4月8日
fix: test_mapping.py #2167
已合并
aijgnem1创建于 2025年4月8日
从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 @pytest.mark.parametrize("dtype", [torch.float16, torch.bfloat16])129 @pytest.mark.parametrize("dtype", [torch.float16, torch.bfloat16])
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 @pytest.mark.parametrize("dtype", [torch.float16, torch.bfloat16])156 @pytest.mark.parametrize("dtype", [torch.float16, torch.bfloat16])
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 @pytest.mark.parametrize("dim", [0, 1, 2, 3])196 @pytest.mark.parametrize("dim", [0, 1, 2, 3])
189 @pytest.mark.parametrize("dtype", [torch.float16, torch.bfloat16])197 @pytest.mark.parametrize("dtype", [torch.float16, torch.bfloat16])
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():