Skip to content

Commit e7c8aba

Browse files
authored
[recipe] fix: Resync Kimi pipeline layout after overrides (#5008)
Signed-off-by: Yu Yao <yaoyu.094@gmail.com>
1 parent 0765d03 commit e7c8aba

6 files changed

Lines changed: 138 additions & 1 deletion

File tree

scripts/training/recipe_runner.py

Lines changed: 28 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -368,6 +368,34 @@ def sync_finetuning_cp_invariants(config: ConfigContainer, *, mode: str) -> Conf
368368
return config
369369

370370

371+
def sync_model_pipeline_layout(
372+
config: ConfigContainer,
373+
*,
374+
cli_overrides: list[str],
375+
) -> ConfigContainer:
376+
"""Rebuild a recipe-owned pipeline layout after PP or VP overrides."""
377+
override_fields = {override.lstrip("+~").split("=", 1)[0] for override in cli_overrides}
378+
topology_fields = {
379+
"model.pipeline_model_parallel_size",
380+
"model.virtual_pipeline_model_parallel_size",
381+
}
382+
if not override_fields.intersection(topology_fields):
383+
return config
384+
if "model.pipeline_model_parallel_layout" in override_fields:
385+
return config
386+
387+
model = getattr(config, "model", None)
388+
layout_builder = getattr(model, "_pipeline_model_parallel_layout_builder", None)
389+
if layout_builder is None:
390+
return config
391+
392+
model.pipeline_model_parallel_layout = layout_builder(
393+
model.pipeline_model_parallel_size,
394+
model.virtual_pipeline_model_parallel_size,
395+
)
396+
return config
397+
398+
371399
def sync_offline_packing_alignment(config: ConfigContainer) -> ConfigContainer:
372400
"""Align offline-packed samples to the resolved length and parallel topology."""
373401
dataset = getattr(config, "dataset", None)

scripts/training/run_recipe.py

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -100,6 +100,7 @@
100100
run_config,
101101
sync_finetuning_cp_invariants,
102102
sync_model_dataset_sequence_length,
103+
sync_model_pipeline_layout,
103104
sync_offline_packing_alignment,
104105
)
105106

@@ -418,6 +419,7 @@ def main(argv: list[str] | None = None) -> None:
418419
recipe = _apply_dataset(recipe, args)
419420
recipe = apply_determinism(recipe, deterministic=args.deterministic)
420421
recipe = apply_cli_overrides(recipe, cli_overrides)
422+
recipe = sync_model_pipeline_layout(recipe, cli_overrides=cli_overrides)
421423
if benchmark_metadata is not None:
422424
recipe = _apply_benchmark_runtime_defaults(recipe, benchmark_metadata, cli_overrides)
423425
configuration_mode = _train_mode(args.mode)

src/megatron/bridge/recipes/kimi/h100/kimi_k2.py

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -26,7 +26,7 @@
2626
from megatron.bridge.training.mixed_precision import MixedPrecisionConfig
2727

2828

29-
def _get_kimi_k2_pipeline_layout(pp_size: int, vp_size: int):
29+
def _get_kimi_k2_pipeline_layout(pp_size: int, vp_size: int | None) -> list[list[str]] | None:
3030
"""Get pipeline layout for Kimi-K2 based on PP and VP size."""
3131
map_pp_vp_to_layout = {
3232
(1, 1): None,
@@ -121,6 +121,7 @@ def kimi_k2_pretrain_512gpu_h100_bf16_config() -> ConfigContainer:
121121
cfg.model.num_layers_in_last_pipeline_stage = None
122122

123123
# Set pipeline layout
124+
cfg.model._pipeline_model_parallel_layout_builder = _get_kimi_k2_pipeline_layout
124125
cfg.model.pipeline_model_parallel_layout = _get_kimi_k2_pipeline_layout(16, 1)
125126

126127
# Tokenizer - uses NullTokenizer with model vocab_size

tests/unit_tests/recipes/kimi/test_kimi_k2.py

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -271,3 +271,4 @@ def test_pretrain_config_pipeline_layout(self):
271271
# Default PP=16, VP=None (1), should have a layout
272272
expected_layout = _get_kimi_k2_pipeline_layout(16, 1)
273273
assert cfg.model.pipeline_model_parallel_layout == expected_layout
274+
assert cfg.model._pipeline_model_parallel_layout_builder is _get_kimi_k2_pipeline_layout

tests/unit_tests/scripts/training/test_recipe_runner.py

Lines changed: 50 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -244,6 +244,56 @@ def test_sync_model_dataset_sequence_length_uses_canonical_dataset_field(recipe_
244244
assert config.model.seq_length == 256
245245

246246

247+
def test_sync_model_pipeline_layout_uses_overridden_topology(recipe_runner: ModuleType) -> None:
248+
layout = [["first"], ["middle"], ["middle"], ["last"]]
249+
layout_builder = Mock(return_value=layout)
250+
config = SimpleNamespace(
251+
model=SimpleNamespace(
252+
pipeline_model_parallel_size=4,
253+
virtual_pipeline_model_parallel_size=1,
254+
pipeline_model_parallel_layout=[["stale"]] * 16,
255+
_pipeline_model_parallel_layout_builder=layout_builder,
256+
)
257+
)
258+
259+
assert (
260+
recipe_runner.sync_model_pipeline_layout(
261+
config,
262+
cli_overrides=[
263+
"model.pipeline_model_parallel_size=4",
264+
"model.virtual_pipeline_model_parallel_size=1",
265+
],
266+
)
267+
is config
268+
)
269+
assert config.model.pipeline_model_parallel_layout == layout
270+
layout_builder.assert_called_once_with(4, 1)
271+
272+
273+
def test_sync_model_pipeline_layout_preserves_explicit_layout_override(recipe_runner: ModuleType) -> None:
274+
layout = [["custom"]]
275+
layout_builder = Mock()
276+
config = SimpleNamespace(
277+
model=SimpleNamespace(
278+
pipeline_model_parallel_size=1,
279+
virtual_pipeline_model_parallel_size=1,
280+
pipeline_model_parallel_layout=layout,
281+
_pipeline_model_parallel_layout_builder=layout_builder,
282+
)
283+
)
284+
285+
recipe_runner.sync_model_pipeline_layout(
286+
config,
287+
cli_overrides=[
288+
"model.pipeline_model_parallel_size=1",
289+
"model.pipeline_model_parallel_layout=[[custom]]",
290+
],
291+
)
292+
293+
assert config.model.pipeline_model_parallel_layout == layout
294+
layout_builder.assert_not_called()
295+
296+
247297
@pytest.mark.parametrize(
248298
"dataset",
249299
[

tests/unit_tests/scripts/training/test_run_recipe.py

Lines changed: 55 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -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

Comments
 (0)