@@ -47,6 +47,7 @@ def _load_module():
4747 "load_recipe" ,
4848 "run_config" ,
4949 "sync_finetuning_cp_invariants" ,
50+ "sync_model_pipeline_layout" ,
5051 "sync_offline_packing_alignment" ,
5152 "sync_model_dataset_sequence_length" ,
5253 ):
@@ -56,6 +57,7 @@ def _load_module():
5657 recipe_runner .apply_runtime_environment .side_effect = lambda config : config
5758 recipe_runner .bootstrap_recipe_environment .side_effect = lambda config , ** _ : config
5859 recipe_runner .sync_finetuning_cp_invariants .side_effect = lambda config , ** _ : config
60+ recipe_runner .sync_model_pipeline_layout .side_effect = lambda config , ** _ : config
5961 recipe_runner .sync_offline_packing_alignment .side_effect = lambda config : config
6062 recipe_runner .sync_model_dataset_sequence_length .side_effect = lambda config : config
6163 recipe_runner .load_forward_step .return_value = object ()
@@ -224,6 +226,59 @@ def test_full_recipe_uses_library_recipe_and_default_llm_step():
224226 handles .recipe_runner .load_forward_step .assert_called_once_with ("llm_step" , mode = "pretrain" )
225227
226228
229+ def test_kimi_supported_pp_vp_override_refreshes_pipeline_layout ():
230+ module , handles = _load_module ()
231+ default_layout = [[f"stage-{ index } " ] for index in range (16 )]
232+ overridden_layout = [[f"stage-{ index } " ] for index in range (4 )]
233+ config = SimpleNamespace (
234+ model = SimpleNamespace (
235+ pipeline_model_parallel_size = 16 ,
236+ virtual_pipeline_model_parallel_size = None ,
237+ pipeline_model_parallel_layout = default_layout ,
238+ _pipeline_model_parallel_layout_builder = lambda pp , vp : overridden_layout
239+ if (pp , vp ) == (4 , 1 )
240+ else default_layout ,
241+ )
242+ )
243+ handles .recipe_runner .load_recipe .return_value = config
244+
245+ def apply_override (current_config , overrides ):
246+ assert overrides == [
247+ "model.pipeline_model_parallel_size=4" ,
248+ "model.virtual_pipeline_model_parallel_size=1" ,
249+ ]
250+ current_config .model .pipeline_model_parallel_size = 4
251+ current_config .model .virtual_pipeline_model_parallel_size = 1
252+ return current_config
253+
254+ def refresh_layout (current_config , * , cli_overrides ):
255+ assert cli_overrides == [
256+ "model.pipeline_model_parallel_size=4" ,
257+ "model.virtual_pipeline_model_parallel_size=1" ,
258+ ]
259+ model = current_config .model
260+ model .pipeline_model_parallel_layout = model ._pipeline_model_parallel_layout_builder (
261+ model .pipeline_model_parallel_size ,
262+ model .virtual_pipeline_model_parallel_size ,
263+ )
264+ return current_config
265+
266+ def validate_layout (* , config , ** _ ):
267+ model = config .model
268+ detected_vp = len (model .pipeline_model_parallel_layout ) // model .pipeline_model_parallel_size
269+ assert detected_vp == model .virtual_pipeline_model_parallel_size , (
270+ "virtual_pipeline_model_parallel_size conflicts with the pipeline layout"
271+ )
272+
273+ handles .recipe_runner .apply_cli_overrides .side_effect = apply_override
274+ handles .recipe_runner .sync_model_pipeline_layout .side_effect = refresh_layout
275+ handles .recipe_runner .run_config .side_effect = validate_layout
276+
277+ module .main (["--recipe" , "kimi_k2_pretrain_config" , "-pp" , "4" , "-vp" , "1" ])
278+
279+ assert config .model .pipeline_model_parallel_layout == overridden_layout
280+
281+
227282@pytest .mark .parametrize (
228283 ("recipe_name" , "mode" , "step_name" ),
229284 [
0 commit comments