Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
19 changes: 18 additions & 1 deletion gpt_builders.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,7 @@

from megatron.core.models.gpt import GPTModel
from megatron.core.models.gpt.experimental_attention_variant_module_specs import (
get_experimental_attention_variant_stage_input_cp_partition_mode,
get_transformer_block_with_experimental_attention_variant_spec,
get_transformer_layer_with_experimental_attention_variant_spec,
)
Expand All @@ -17,6 +18,7 @@
get_gpt_heterogeneous_layer_spec,
)
from megatron.core.transformer.spec_utils import import_module
from megatron.core.utils import get_pg_rank
from megatron.training import get_args, print_rank_0
from megatron.training.arguments import core_transformer_config_from_args
from megatron.training.yaml_arguments import core_transformer_config_from_yaml
Expand All @@ -29,14 +31,28 @@ def gpt_builder(args, pre_process, post_process, vp_stage=None, config=None, pg_
config = core_transformer_config_from_yaml(args, "language_model")
else:
config = core_transformer_config_from_args(args)
cp_stage_entry_partition_mode = (
"zigzag" if config.cp_partition_mode == "auto" else config.cp_partition_mode
)
if args.spec is not None:
transformer_layer_spec = import_module(args.spec)
else:
use_te = args.transformer_impl == "transformer_engine"

if args.experimental_attention_variant is not None:
pp_rank = (
get_pg_rank(pg_collection.pp)
if pg_collection is not None and hasattr(pg_collection, "pp")
else None
)
if config.cp_partition_mode == "auto":
cp_stage_entry_partition_mode = (
get_experimental_attention_variant_stage_input_cp_partition_mode(
config=config, vp_stage=vp_stage, pp_rank=pp_rank
)
)
transformer_layer_spec = get_transformer_block_with_experimental_attention_variant_spec(
config=config, vp_stage=vp_stage
config=config, vp_stage=vp_stage, pp_rank=pp_rank
)
elif args.num_experts:
# Define the decoder block spec
Expand Down Expand Up @@ -107,6 +123,7 @@ def gpt_builder(args, pre_process, post_process, vp_stage=None, config=None, pg_
mtp_block_spec=mtp_block_spec,
vp_stage=vp_stage,
pg_collection=pg_collection,
cp_stage_entry_partition_mode=cp_stage_entry_partition_mode,
)

return model
Expand Down
15 changes: 15 additions & 0 deletions hybrid_builders.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,9 @@
# Copyright (c) 2025-2026, NVIDIA CORPORATION. All rights reserved.

from megatron.core.models.hybrid.hybrid_layer_specs import hybrid_inference_stack_spec
from megatron.core.models.hybrid.hybrid_layer_allocation import (

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

[SUGGESTION Naming] Import order: hybrid_layer_allocation sorts before hybrid_layer_specs alphabetically, so isort will want this new import placed above the hybrid_layer_specs line. Per CLAUDE.md ("after editing imports … run uv run isort on those files"), please run isort here to fix ordering before merge.

get_hybrid_stage_input_cp_partition_mode_for_stage,
)
from megatron.core.models.hybrid.hybrid_model import HybridModel
from megatron.core.transformer import MLATransformerConfig, TransformerConfig
from megatron.core.transformer.spec_utils import ModuleSpec, import_module
Expand Down Expand Up @@ -75,6 +78,17 @@ def hybrid_builder(args, pre_process, post_process, vp_stage=None, config=None,
else:
raise ValueError("You must provide a valid hybrid layer spec via --spec")

cp_stage_entry_partition_mode = config.cp_partition_mode
if config.cp_partition_mode == "auto":
cp_stage_entry_partition_mode = get_hybrid_stage_input_cp_partition_mode_for_stage(
config,
args.hybrid_layer_pattern,
getattr(pg_collection, "pp", None),
vp_stage,
first_stage_layers=config.num_layers_in_first_pipeline_stage,
last_stage_layers=config.num_layers_in_last_pipeline_stage,
)

model = HybridModel(
config=config,
hybrid_stack_spec=hybrid_stack_spec,
Expand All @@ -91,6 +105,7 @@ def hybrid_builder(args, pre_process, post_process, vp_stage=None, config=None,
rotary_base=args.rotary_base,
pg_collection=pg_collection,
vp_stage=vp_stage,
cp_stage_entry_partition_mode=cp_stage_entry_partition_mode,
)

for l in range(model.decoder.num_layers_per_pipeline_rank):
Expand Down
307 changes: 0 additions & 307 deletions megatron/core/context_parallel_layout.py

This file was deleted.

Loading