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
2 changes: 2 additions & 0 deletions src/megatron/bridge/recipes/qwen/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -21,6 +21,7 @@
from .gb200.qwen35 import (
qwen35_text_9b_pretrain_8gpu_gb200_bf16_config,
qwen35_text_35b_a3b_pretrain_8gpu_gb200_bf16_config,
qwen35_text_35b_a3b_sft_long_context_16gpu_gb200_fp8mx_config,
)

# Qwen2 models
Expand Down Expand Up @@ -175,4 +176,5 @@
"qwen35_text_9b_pretrain_8gpu_gb200_bf16_config",
"qwen35_text_35b_a3b_pretrain_config",
"qwen35_text_35b_a3b_pretrain_8gpu_gb200_bf16_config",
"qwen35_text_35b_a3b_sft_long_context_16gpu_gb200_fp8mx_config",
]
2 changes: 2 additions & 0 deletions src/megatron/bridge/recipes/qwen/gb200/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -18,11 +18,13 @@
from megatron.bridge.recipes.qwen.gb200.qwen35 import (
qwen35_text_9b_pretrain_8gpu_gb200_bf16_config,
qwen35_text_35b_a3b_pretrain_8gpu_gb200_bf16_config,
qwen35_text_35b_a3b_sft_long_context_16gpu_gb200_fp8mx_config,
)


__all__ = [
"qwen3_30b_a3b_pretrain_8gpu_gb200_fp8mx_config",
"qwen35_text_9b_pretrain_8gpu_gb200_bf16_config",
"qwen35_text_35b_a3b_pretrain_8gpu_gb200_bf16_config",
"qwen35_text_35b_a3b_sft_long_context_16gpu_gb200_fp8mx_config",
]
123 changes: 120 additions & 3 deletions src/megatron/bridge/recipes/qwen/gb200/qwen35.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,22 +12,139 @@
# See the License for the specific language governing permissions and
# limitations under the License.

"""GB200 text-only pretraining recipes for Qwen3.5 dense and MoE models."""
"""GB200 text-only pretraining and SFT recipes for Qwen3.5 models."""

from __future__ import annotations

from copy import deepcopy

import torch
from transformers import AutoConfig

from megatron.bridge import AutoBridge
from megatron.bridge.recipes.common import _pretrain_common
from megatron.bridge.recipes.common import _pretrain_common, _sft_common
from megatron.bridge.recipes.utils.dataset_utils import default_coderforge_config
from megatron.bridge.recipes.utils.environment_utils import COMMON_RECIPE_ENV_VARS
from megatron.bridge.recipes.utils.optimizer_utils import distributed_fused_adam_with_cosine_annealing
from megatron.bridge.training.comm_overlap import CommOverlapConfig
from megatron.bridge.training.config import ConfigContainer
from megatron.bridge.training.mixed_precision import bf16_mixed
from megatron.bridge.training.mixed_precision import bf16_mixed, bf16_with_mxfp8_mixed


_QWEN35_9B_BASE = "Qwen/Qwen3.5-9B-Base"
_QWEN35_35B_A3B_BASE = "Qwen/Qwen3.5-35B-A3B-Base"
_QWEN35_35B_A3B_INSTRUCT_REVISION = "59d61f3ce65a6d9863b86d2e96597125219dc754" # pragma: allowlist secret
_CODERFORGE_REVISION = "060fca96cf723b2ebab3181e9e59fafd273df3cb" # pragma: allowlist secret


def qwen35_text_35b_a3b_sft_long_context_16gpu_gb200_fp8mx_config(
seq_length: int = 131072,
*,
hf_path: str = "Qwen/Qwen3.5-35B-A3B",
hf_revision: str | None = _QWEN35_35B_A3B_INSTRUCT_REVISION,
) -> ConfigContainer:
"""MXFP8 text SFT with packed chat data, TP1/CP8/EP16 and one MTP layer.

Set a pretrained text checkpoint before training. The default CoderForge
source is normalized and packed by the dataset builder. Model config and
tokenizer share a pinned revision; use ``hf_revision=None`` for an extracted
local causal-LM checkpoint. No vision model is constructed.

Args:
seq_length: Packed token budget, a positive multiple of 16 for CP8.
hf_path: HF model ID or extracted local text checkpoint.
hf_revision: Matching model/tokenizer revision, or None for local inputs.

Returns:
Training configuration with MXFP8 parameter gather and CuTeDSL grouped MLP.
"""
if seq_length <= 0 or seq_length % 16:
raise ValueError("seq_length must be positive and divisible by 16 for CP8.")
cfg = _sft_common()
hf_config = AutoConfig.from_pretrained(hf_path, revision=hf_revision)
text_config = deepcopy(getattr(hf_config, "text_config", hf_config))
if text_config.model_type != "qwen3_5_moe_text":
raise ValueError("Expected a Qwen3.5 MoE text configuration.")
if seq_length > text_config.max_position_embeddings:
raise ValueError("seq_length exceeds the text model's max_position_embeddings.")
if hasattr(hf_config, "text_config"):
text_config.tie_word_embeddings = hf_config.tie_word_embeddings
text_config.architectures = ["Qwen3_5MoeForCausalLM"]
cfg.model = AutoBridge.from_hf_config(text_config).to_megatron_provider(load_weights=False)
cfg.tokenizer.tokenizer_model = hf_path
cfg.tokenizer.hf_tokenizer_kwargs = {"revision": hf_revision}

cfg.model.seq_length = seq_length
cfg.model.pipeline_dtype = torch.bfloat16
cfg.model.context_parallel_size = 8
cfg.model.expert_model_parallel_size = 16
cfg.model.expert_tensor_parallel_size = 1
cfg.model.mtp_num_layers = 1
cfg.model.bias_activation_fusion = True
cfg.model.cross_entropy_fusion_impl = "te"
cfg.model.calculate_per_token_loss = True
cfg.model.recompute_granularity = "selective"
cfg.model.recompute_modules = ["gdn_norm_out", "moe"]
cfg.model.moe_token_dispatcher_type = "flex"
cfg.model.moe_flex_dispatcher_backend = "hybridep"
cfg.model.moe_flex_dispatcher_num_sms = 32
cfg.model.moe_hybridep_pad_uneven_dispatch_inputs = True
cfg.model.moe_router_fusion = True
cfg.model.moe_use_grouped_tensor = True
cfg.model.use_transformer_engine_op_fuser = True
cfg.model.moe_mlp_glu_interleave_size = 32
cfg.model.moe_single_grouped_weight = True

# Pad each runtime segment to 2*CP, not each conversation to seq_length.
# Physical and logical cu_seqlens remain distinct for attention and MTP.
cfg.dataset = default_coderforge_config(
seq_length=seq_length,
enable_offline_packing=True,
pad_seq_to_mult=16,
)
cfg.dataset.dataset_kwargs = {"pad_to_max_length": True}
cfg.dataset.hf_dataset.load_kwargs = {"revision": _CODERFORGE_REVISION}
cfg.dataset.do_validation = True
cfg.dataset.hf_validation_proportion = 0.05
# Preserve the original long-context recipe's sampling and worker lifecycle.
cfg.dataset.seed = 1234
cfg.dataset.persistent_workers = True
cfg.train.train_iters = 500
cfg.train.global_batch_size = 32
cfg.train.manual_gc = True
cfg.train.manual_gc_interval = 100
cfg.train.manual_gc_eval = 100
cfg.rng.seed = 1234
cfg.validation.eval_interval = 50
cfg.validation.eval_iters = 8
# Preserve the VL baseline's LR schedule rather than silently retuning SFT.
cfg.optimizer, cfg.scheduler = distributed_fused_adam_with_cosine_annealing(
lr_warmup_iters=200, lr_decay_iters=300000, max_lr=2e-5, min_lr=2e-6
)
cfg.optimizer.overlap_param_gather = True
cfg.mixed_precision = bf16_with_mxfp8_mixed()
cfg.checkpoint.load_main_params_from_ckpt = True
cfg.checkpoint.load_optim = False
cfg.checkpoint.load_rng = False
cfg.ddp.use_distributed_optimizer = True
cfg.ddp.data_parallel_sharding_strategy = "optim_grads_params"
cfg.ddp.overlap_grad_reduce = True
cfg.ddp.overlap_param_gather = True
cfg.env_vars = {
**COMMON_RECIPE_ENV_VARS,
"CUDA_DEVICE_MAX_CONNECTIONS": 1,
"NVTE_BWD_LAYERNORM_SM_MARGIN": 20,
"NVTE_FWD_LAYERNORM_SM_MARGIN": 20,
"NUM_OF_HYBRID_EP_RANKS_PER_NVLINK_DOMAIN": 16,
"NUM_OF_TOKENS_PER_CHUNK_COMBINE_API": 128,
"NVLINK_DOMAIN_SIZE": 72,
"USE_MNNVL": 1,
"NVTE_CUTEDSL_FUSED_GROUPED_MLP": 1,
"CUDNN_FE_GROUPED_GEMM_DYNAMIC_MNKL": 1,
"NVTE_GROUPED_LINEAR_SINGLE_PARAM": 1,
}
cfg.checkpoint.load = None
return cfg


def qwen35_text_9b_pretrain_8gpu_gb200_bf16_config() -> ConfigContainer:
Expand Down
9 changes: 8 additions & 1 deletion tests/unit_tests/recipes/recipe_test_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -261,7 +261,14 @@ def patch_recipe_construction_dependencies(monkeypatch: pytest.MonkeyPatch) -> N

def load_offline_auto_config(*args: object, **kwargs: object) -> SimpleNamespace:
del args, kwargs
return SimpleNamespace(text_config=SimpleNamespace(architectures=None))
return SimpleNamespace(
tie_word_embeddings=False,
text_config=SimpleNamespace(
architectures=None,
model_type="qwen3_5_moe_text",
max_position_embeddings=262144,
),
)

monkeypatch.setattr(
AutoConfig,
Expand Down
158 changes: 158 additions & 0 deletions tests/unit_tests/recipes/test_qwen35_text_long_context.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,158 @@
# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved.

import importlib
from copy import deepcopy
from pathlib import Path
from types import SimpleNamespace

import pytest
import torch
from transformers import AutoConfig

from megatron.bridge.recipes.qwen.gb200 import qwen35


pytestmark = pytest.mark.unit
_MXFP8 = qwen35.qwen35_text_35b_a3b_sft_long_context_16gpu_gb200_fp8mx_config


@pytest.fixture
def hf_config(monkeypatch):
config = AutoConfig.for_model(
"qwen3_5_moe",
text_config={
"num_hidden_layers": 4,
"hidden_size": 512,
"num_attention_heads": 4,
"num_key_value_heads": 4,
"head_dim": 128,
"max_position_embeddings": 262144,
"num_experts": 32,
"num_experts_per_tok": 2,
"moe_intermediate_size": 128,
"shared_expert_intermediate_size": 128,
"linear_num_key_heads": 8,
"linear_num_value_heads": 8,
"linear_key_head_dim": 128,
"linear_value_head_dim": 128,
"tie_word_embeddings": True,
},
tie_word_embeddings=False,
)
calls = []

def load(path, **kwargs):
calls.append((path, kwargs))
return config

monkeypatch.setattr(qwen35.AutoConfig, "from_pretrained", load)
return config, calls


def test_text_long_context_validates_and_preserves_inputs(hf_config, monkeypatch):
config, calls = hf_config
original = deepcopy(config.to_dict())
training_module = importlib.import_module("megatron.bridge.training.config")
monkeypatch.setattr(training_module, "get_world_size_safe", lambda: 16)
monkeypatch.setattr(
torch.cuda, "get_device_properties", lambda index: SimpleNamespace(major=10, name="NVIDIA GB200")
)

cfg = _MXFP8()
cfg.validate()

assert config.to_dict() == original
assert calls == [("Qwen/Qwen3.5-35B-A3B", {"revision": qwen35._QWEN35_35B_A3B_INSTRUCT_REVISION})]
assert cfg.tokenizer.hf_tokenizer_kwargs["revision"] == calls[0][1]["revision"]
assert cfg.model.share_embeddings_and_output_weights is False
assert cfg.model.mtp_num_layers == 1
assert cfg.model.mtp_loss_scaling_factor == 0.1
assert cfg.model.context_parallel_size == 8
assert cfg.model.expert_model_parallel_size == 16
assert cfg.model.tensor_model_parallel_size == 1
assert cfg.model.sequence_parallel is False
assert cfg.model.recompute_modules == ["gdn_norm_out", "moe"]
assert cfg.dataset.hf_dataset.dataset_name == "coderforge"
assert cfg.dataset.hf_dataset.load_kwargs == {"revision": qwen35._CODERFORGE_REVISION}
assert cfg.dataset.dataset_root is None
assert cfg.dataset.hf_validation_proportion == 0.05
assert cfg.dataset.offline_packing_specs.pad_seq_to_mult == 16
assert cfg.dataset.offline_packing_specs.packed_sequence_size == 131072
assert cfg.dataset.dataset_kwargs["pad_to_max_length"] is True
assert cfg.dataset.seed == 1234
assert cfg.dataset.persistent_workers is True
assert cfg.train.global_batch_size == 32
assert cfg.train.micro_batch_size == 1
assert cfg.train.train_iters == 500
assert cfg.scheduler.lr_warmup_iters == 200
assert cfg.scheduler.lr_decay_iters == 300000
assert cfg.ddp.grad_reduce_in_fp32 is True


def test_verified_mxfp8_kernel_and_parameter_settings(hf_config):
cfg = _MXFP8()
assert cfg.mixed_precision.fp8_recipe == "mxfp8"
assert cfg.mixed_precision.fp8_param_gather is True
assert cfg.mixed_precision.reuse_grad_buf_for_mxfp8_param_ag is True
assert cfg.checkpoint.load_main_params_from_ckpt is True
assert cfg.checkpoint.load_optim is False
assert cfg.checkpoint.load_rng is False
assert cfg.model.moe_use_grouped_tensor is True
assert cfg.model.use_transformer_engine_op_fuser is True
assert cfg.model.moe_mlp_glu_interleave_size == 32
assert cfg.model.moe_single_grouped_weight is True
assert cfg.model.moe_single_grouped_bias is False
assert cfg.env_vars["NUM_OF_HYBRID_EP_RANKS_PER_NVLINK_DOMAIN"] == 16
assert cfg.env_vars["NVTE_CUTEDSL_FUSED_GROUPED_MLP"] == 1
assert cfg.env_vars["CUDNN_FE_GROUPED_GEMM_DYNAMIC_MNKL"] == 1
assert cfg.env_vars["NVTE_GROUPED_LINEAR_SINGLE_PARAM"] == 1


@pytest.mark.parametrize("seq_length", [0, -16, 15, 131073])
def test_invalid_sequence_length_fails_before_hub_access(hf_config, seq_length):
_, calls = hf_config
with pytest.raises(ValueError, match="positive and divisible by 16"):
_MXFP8(seq_length)
assert calls == []


def test_local_text_config_uses_matching_tokenizer(hf_config, monkeypatch):
config, _ = hf_config
monkeypatch.setattr(qwen35.AutoConfig, "from_pretrained", lambda *_args, **_kwargs: config.text_config)

cfg = _MXFP8(1024, hf_path="local-text-checkpoint", hf_revision=None)

assert cfg.tokenizer.tokenizer_model == "local-text-checkpoint"
assert cfg.tokenizer.hf_tokenizer_kwargs == {"revision": None}
assert cfg.model.seq_length == 1024
assert cfg.dataset.offline_packing_specs.packed_sequence_size == 1024


def test_wrong_model_type_is_rejected(hf_config):
config, _ = hf_config
config.text_config.model_type = "qwen3_5_text"
with pytest.raises(ValueError, match="Expected a Qwen3.5 MoE"):
_MXFP8()


@pytest.mark.parametrize("seq_length", [16, 32768, 65536, 131072])
def test_aligned_context_lengths(hf_config, seq_length):
cfg = _MXFP8(seq_length)
assert cfg.model.seq_length == cfg.dataset.seq_length == seq_length
assert cfg.dataset.offline_packing_specs.packed_sequence_size == seq_length


def test_context_exceeding_model_limit_is_rejected(hf_config):
config, _ = hf_config
with pytest.raises(ValueError, match="max_position_embeddings"):
_MXFP8(config.text_config.max_position_embeddings + 16)


def test_public_runner_finds_unique_text_recipe(hf_config, monkeypatch):
monkeypatch.syspath_prepend(str(Path(__file__).resolve().parents[3] / "scripts/training"))
runner = importlib.import_module("recipe_runner")
assert runner.find_library_recipe(_MXFP8.__name__) is _MXFP8
assert runner.find_benchmark_recipe(_MXFP8.__name__) is None
cfg = runner.load_recipe(_MXFP8.__name__)
assert cfg.model.mtp_num_layers == 1
assert cfg.dataset.hf_dataset.dataset_name == "coderforge"
7 changes: 5 additions & 2 deletions tests/unit_tests/recipes/test_qwen_recipes.py
Original file line number Diff line number Diff line change
Expand Up @@ -97,16 +97,19 @@ def to_megatron_provider(self, load_weights: bool = False):

class _FakeTextConfig:
architectures = None
model_type = "qwen3_5_moe_text"
max_position_embeddings = 262144


class _FakeRootConfig:
text_config = _FakeTextConfig()
tie_word_embeddings = False


class _FakeAutoConfig:
@staticmethod
def from_pretrained(hf_path: str):
# Ignore hf_path; return a unified config with a nested text config.
def from_pretrained(hf_path: str, *, revision: str | None = None):
# Match the pinned-config API without downloading a model configuration.
return _FakeRootConfig()


Expand Down
Loading