Skip to content
Merged
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
15 changes: 15 additions & 0 deletions src/megatron/bridge/recipes/qwen/qwen3_next.py
Original file line number Diff line number Diff line change
Expand Up @@ -28,6 +28,7 @@
from megatron.bridge.training.config import (
CheckpointConfig,
ConfigContainer,
DistributedInitConfig,
FinetuningDatasetConfig,
GPTDatasetConfig,
LoggerConfig,
Expand Down Expand Up @@ -86,6 +87,7 @@ class Qwen3NextCommonKwargs(TypedDict, total=False):
comm_overlap_config: CommOverlapConfig | None
# Performance optimization knobs
enable_deepep: bool
disable_jit_fuser: bool


class Qwen3NextFinetuneKwargs(Qwen3NextCommonKwargs, total=False):
Expand Down Expand Up @@ -173,6 +175,7 @@ def _qwen3_next_common(
precision_config: MixedPrecisionConfig | str | None = None,
comm_overlap_config: CommOverlapConfig | None = None,
enable_deepep: bool = False,
disable_jit_fuser: bool | None = None,
) -> ConfigContainer:
"""
Create a pre-training configuration for Qwen3-Next models using a given HuggingFace path.
Expand Down Expand Up @@ -213,6 +216,7 @@ def _qwen3_next_common(
precision_config (MixedPrecisionConfig | str | None): Precision configuration for the model.
comm_overlap_config (CommOverlapConfig | None): Communication overlap configuration.
enable_deepep (bool): Whether to enable DEEPEP for MoE.
disable_jit_fuser (bool): Whether to disable the JIT fuser. Necessary for Qwen3-Next to work on Blackwell.

Returns:
ConfigContainer: Configuration for pre-training.
Expand Down Expand Up @@ -277,6 +281,10 @@ def _qwen3_next_common(
)
scheduler.no_weight_decay_cond_type = "qwen3_next"

# If user does not specify, check if we are on Blackwell.
if disable_jit_fuser is None:
disable_jit_fuser = torch.cuda.get_device_properties(0).major == 10

# Config Container
cfg = ConfigContainer(
model=model_cfg,
Expand All @@ -292,6 +300,7 @@ def _qwen3_next_common(
),
optimizer=opt_config,
scheduler=scheduler,
dist=DistributedInitConfig(disable_jit_fuser=disable_jit_fuser),
ddp=DistributedDataParallelConfig(
check_for_nan_in_grad=True,
grad_reduce_in_fp32=True,
Expand Down Expand Up @@ -421,6 +430,7 @@ def _qwen3_next_finetune_common(
precision_config: MixedPrecisionConfig | str | None = "bf16_mixed",
comm_overlap_config: CommOverlapConfig | None = None,
enable_deepep: bool = False,
disable_jit_fuser: bool | None = None,
) -> ConfigContainer:
"""Common finetuning configuration for Qwen3-Next model."""

Expand Down Expand Up @@ -508,6 +518,10 @@ def _qwen3_next_finetune_common(
tokenizer_model=hf_path,
)

# If user does not specify, check if we are on Blackwell.
if disable_jit_fuser is None:
disable_jit_fuser = torch.cuda.get_device_properties(0).major == 10

return ConfigContainer(
model=model_cfg,
train=TrainingConfig(
Expand All @@ -522,6 +536,7 @@ def _qwen3_next_finetune_common(
),
optimizer=opt_cfg,
scheduler=scheduler_cfg,
dist=DistributedInitConfig(disable_jit_fuser=disable_jit_fuser),
ddp=DistributedDataParallelConfig(
check_for_nan_in_grad=True,
grad_reduce_in_fp32=True,
Expand Down
3 changes: 3 additions & 0 deletions src/megatron/bridge/training/config.py
Original file line number Diff line number Diff line change
Expand Up @@ -188,6 +188,9 @@ class DistributedInitConfig:
distributed_timeout_seconds_after_init: int | None = None
"""Timeout in seconds for process groups after initialization. This timeout is applied to all process groups after initialization and the first iteration completes."""

disable_jit_fuser: bool = False
"""Disable the JIT fuser."""


@dataclass
class RerunStateMachineConfig:
Expand Down
6 changes: 6 additions & 0 deletions src/megatron/bridge/training/setup.py
Original file line number Diff line number Diff line change
Expand Up @@ -22,6 +22,7 @@
from megatron.core.config import set_experimental_flag
from megatron.core.distributed import DistributedDataParallel, DistributedDataParallelConfig, finalize_model_grads
from megatron.core.distributed.fsdp.mcore_fsdp_adapter import FullyShardedDataParallel as megatron_FSDP
from megatron.core.jit import disable_jit_fuser
from megatron.core.optimizer import MegatronOptimizer
from megatron.core.optimizer_param_scheduler import OptimizerParamScheduler
from megatron.core.rerun_state_machine import RerunDataIterator
Expand Down Expand Up @@ -114,6 +115,11 @@ def setup(
# Conditionally enable experimental features for Megatron Core
set_experimental_flag(cfg.dist.enable_megatron_core_experimental)

# Disable the JIT fuser if requested
if cfg.dist.disable_jit_fuser:
print_rank_0("Disabling JIT fuser.")
disable_jit_fuser()

# Initialize async checkpoint worker if enabled (idempotent if already initialized)
state.initialize_async_checkpoint_worker()

Expand Down