Add HybridModel MBridge support for nemo:26.08 (#2005)

### Description

MBridge main (nemo:26.08) will initialize Nemotron-H as HybridModel
instead of MambaModel (subclassed of HybridModel). Also make minimum
nemo container 26.04

### Testing 

Tested Nemotron-3-Nano PTQ with MBridge main (fails otherwise)

Tested locally `tests/gpu_megatron` and `tests/examples/megatron_bridge`
with `nemo:26.06.01` + Mount latest MBridge/Mcore

GH CICD tests will be added with nemo:26.08 release

<!-- This is an auto-generated comment: release notes by coderabbit.ai
-->
## Summary by CodeRabbit

* **New Features**
* Added TE HybridModel stack-spec support and enabled Megatron-Bridge
hybrid providers/models for export/import and runtime handling.
* **Updates**
* Dataset packing now oversamples raw text at **16x** and improves the
packed-mode underflow warning.
* Quantization: `--quant_cfg` now defaults to `None` unless explicitly
set (or via `--recipe`).
* Distillation example: validation settings are provided via a dedicated
top-level validation configuration.
* Improved plugin import warnings to report the originating call
location; model stats now support HybridModel.
* **Deprecations**
* Megatron-Bridge / Megatron-LM optimization features now require NeMo
container `nemo:26.04` or newer (`nemo:26.06` recommended).
  * The Mamba stack specification helper is deprecated.
<!-- end of auto-generated comment: release notes by coderabbit.ai -->

---------

Signed-off-by: Keval Morabia <28916987+kevalmorabia97@users.noreply.github.com>
This commit is contained in:
Keval Morabia
2026-07-23 18:26:15 +05:30
committed by GitHub
parent 3c9137d28e
commit 01c708e792
19 changed files with 195 additions and 102 deletions
+1
View File
@@ -22,6 +22,7 @@ Experimental
- Renamed ``examples/llm_ptq`` to ``examples/hf_ptq`` to reflect that it covers Hugging Face LLM **and** VLM PTQ. A relative symlink ``examples/llm_ptq`` -> ``hf_ptq`` keeps existing paths and commands working; it will be removed in a future release. Please update references to the new ``examples/hf_ptq`` path.
- Consolidated ``examples/vlm_ptq`` into ``examples/hf_ptq``. Vision-language model PTQ now shares the ``hf_ptq.py`` entry point and ``scripts/huggingface_example.sh``; pass ``--vlm`` to run the TensorRT-LLM multimodal quickstart smoke test. The ``examples/vlm_ptq/scripts/huggingface_example.sh`` entry point is deprecated: it now prints a warning and forwards to the ``hf_ptq`` script with ``--vlm``, and will be removed in a future release. See `examples/hf_ptq/README.md <https://github.com/NVIDIA/Model-Optimizer/tree/main/examples/hf_ptq#vlm-quantization>`__.
- Dropped VILA / NVILA vision-language model support in ``examples/hf_ptq``. VILA's modeling code requires ``transformers<=4.50.0``, which conflicts with ModelOpt's minimum supported ``transformers`` version. The VILA-specific bootstrap (repo clone, ``requirements-vila.txt``) and loading paths in ``example_utils.py`` have been removed.
- Bump minimum nemo container requirement to ``nemo:26.04`` (recommended ``nemo:26.06``) for Megatron-Bridge / Megatron-LM optimization features.
**New Features**
+2 -4
View File
@@ -40,6 +40,7 @@ from megatron.bridge.training.config import (
RNGConfig,
TokenizerConfig,
TrainingConfig,
ValidationConfig,
)
from megatron.bridge.training.distill import distill
from megatron.bridge.training.post_training.checkpointing import has_modelopt_state
@@ -325,15 +326,12 @@ def main(args: argparse.Namespace):
model=distill_provider,
train=TrainingConfig(
train_iters=args.train_iters,
eval_interval=args.eval_interval,
eval_iters=args.eval_iters,
global_batch_size=args.gbs,
micro_batch_size=args.mbs,
manual_gc=True,
manual_gc_interval=100,
),
# TODO: Replace validation args in train with validation config once we drop nemo:26.02 container support
# validation=ValidationConfig(eval_interval=args.eval_interval, eval_iters=args.eval_iters),
validation=ValidationConfig(eval_iters=args.eval_iters, eval_interval=args.eval_interval),
optimizer=optimizer_config,
scheduler=scheduler_config,
ddp=DistributedDataParallelConfig(
+15 -2
View File
@@ -46,6 +46,15 @@ import re
import torch
from megatron.bridge import AutoBridge
from megatron.bridge.models.mamba.mamba_provider import MambaModelProvider
try: # nemo:26.08+
from megatron.bridge.models.hybrid.hybrid_provider import HybridModelProvider
# MambaModelProvider subclasses HybridModelProvider on nemo:26.08+, so the tuple covers both.
_HYBRID_PROVIDER_TYPES: tuple[type, ...] = (MambaModelProvider, HybridModelProvider)
except ImportError: # nemo:26.06 and earlier
_HYBRID_PROVIDER_TYPES = (MambaModelProvider,)
from transformers import (
AutoConfig,
AutoModelForCausalLM,
@@ -541,7 +550,7 @@ def main(args: argparse.Namespace):
mto.ModeloptStateManager.remove_state(language_model)
if is_vlm:
_log_vlm_param_breakdown(unwrapped_model, language_model, "after pruning")
if isinstance(provider, MambaModelProvider):
if isinstance(provider, _HYBRID_PROVIDER_TYPES):
hybrid_key = (
"hybrid_override_pattern"
if hasattr(unwrapped_model, "hybrid_override_pattern")
@@ -615,6 +624,8 @@ def main(args: argparse.Namespace):
if sorted_layers is not None
else set(range(1, mcore_cfg.num_layers + 1))
)
# layer_types is the HF per-layer attention-cadence field (mcore's linear_attention_freq /
# moe_layer_freq have no HF equivalent under those names, so only layer_types needs slicing).
if hasattr(text_cfg, "layer_types"):
text_cfg.layer_types = [
lt for i, lt in enumerate(text_cfg.layer_types) if i + 1 in kept_layer_nums
@@ -635,7 +646,9 @@ def main(args: argparse.Namespace):
"distillation cannot recover this vision-path change -- consider full VLM "
"training/distillation instead of LM-only to recover vision quality."
)
if isinstance(provider, MambaModelProvider) and hasattr(hf_cfg, "hybrid_override_pattern"):
if isinstance(provider, _HYBRID_PROVIDER_TYPES) and hasattr(
hf_cfg, "hybrid_override_pattern"
):
hf_cfg.hybrid_override_pattern = getattr(unwrapped_model, hybrid_key)
text_cfg.num_hidden_layers = mcore_cfg.num_layers
# Mark MTP as disabled on the HF text config written after pruning
+1 -1
View File
@@ -124,7 +124,7 @@ def get_args() -> argparse.Namespace:
parser.add_argument(
"--quant_cfg",
type=str,
default="fp8",
default=None,
help=(
f"Quantization config. Preset names / short aliases: {', '.join(QUANT_CFG_CHOICES)}. "
"You can also pass any full config name exposed by modelopt (e.g. FP8_DEFAULT_CFG). "
+40 -9
View File
@@ -88,6 +88,9 @@ except ImportError:
# returns the first match in insertion order: MambaModel is registered first, so
# MambaModel instances dispatch to MambaModel whether or not MambaModel overrides forward.
try:
from megatron.core.models.hybrid.hybrid_layer_specs import (
hybrid_stack_spec as _te_hybrid_stack_spec,
)
from megatron.core.models.hybrid.hybrid_model import HybridModel
SUPPORTED_MODELS[HybridModel] = "megatron.core.models.hybrid.HybridModel"
@@ -99,11 +102,11 @@ except ImportError:
# Attention module types that _DynamicTransformerLayer converts.
_ATTENTION_TYPES: tuple[type, ...] = (SelfAttention, MLASelfAttention, GatedDeltaNet)
__all__ = ["get_te_mamba_stack_spec"]
__all__ = ["get_te_hybrid_stack_spec", "get_te_mamba_stack_spec"]
def get_te_mamba_stack_spec(moe_grouped_gemm: bool = False) -> ModuleSpec:
"""Return the TE Mamba stack spec."""
"""[Deprecated] Return the TE Mamba stack spec."""
assert HAS_MAMBA
if moe_grouped_gemm:
return _te_mamba_stack_spec
@@ -118,6 +121,21 @@ def get_te_mamba_stack_spec(moe_grouped_gemm: bool = False) -> ModuleSpec:
return te_mamba_stack_spec
def get_te_hybrid_stack_spec(moe_grouped_gemm: bool = False) -> ModuleSpec:
"""Return the TE Hybrid stack spec."""
assert HAS_HYBRID
if moe_grouped_gemm:
return _te_hybrid_stack_spec
# The upstream TE hybrid stack spec hardcodes TEGroupedMLP for MoE.
# Replace it with SequentialMLP (TE linear layers, no grouped gemm dependency).
te_hybrid_stack_spec = copy.deepcopy(_te_hybrid_stack_spec)
te_hybrid_stack_spec.submodules.moe_layer.submodules.mlp = get_moe_module_spec(
use_te=True, num_experts=8, moe_grouped_gemm=False
)
return te_hybrid_stack_spec
# Local Parallel Linear DynamicModules ##########################################################################
class _DynamicParallelLinear(DynamicModule):
"""A parallel linear layer with dynamic hyperparams."""
@@ -1055,8 +1073,12 @@ class _MambaContextParallelProxy:
return mixer.d_inner
if name in ("nheads_local_tp", "nheads_local_tpcp"):
return mixer.nheads
if name == "conv1d_cp1":
if name == "conv1d_cp1": # nemo:26.06 and earlier: conv is a module
return mixer.conv1d
if name == "conv1d_weight_cp1": # nemo:26.08+: raw conv parameters (dynamically sliced)
return mixer.conv1d_weight
if name == "conv1d_bias_cp1": # nemo:26.08+
return mixer.conv1d_bias
if name == "dt_bias_cp1":
return mixer.dt_bias
if name == "A_log_cp1":
@@ -1122,11 +1144,19 @@ class _DynamicMambaMixer(DynamicModule):
DMRegistry.convert(self.in_proj, input_size=hidden_size, output_size=in_proj_output_size)
conv_dim = build_concat_hp([d_inner, bc]) # z, B, C
DMRegistry.convert(self.conv1d)
self.conv1d.in_channels = conv_dim
self.conv1d.out_channels = conv_dim
ks = self.conv1d.get_hparam("kernel_size")
ks.choices = [ks.original]
if hasattr(self, "conv1d"): # nemo:26.06 and earlier: a depthwise `nn.Conv1d` module.
DMRegistry.convert(self.conv1d)
self.conv1d.in_channels = conv_dim
self.conv1d.out_channels = conv_dim
ks = self.conv1d.get_hparam("kernel_size")
ks.choices = [ks.original]
else: # nemo:26.08+: the conv is stored as raw parameters
def _slice_conv(mod, val, _hp=conv_dim):
return get_sliced_tensor_by_slices(val, [_hp.active_slice])
self._register_dynamic_attribute("conv1d_weight", _slice_conv) # [conv_dim, 1, d_conv]
self._register_dynamic_attribute("conv1d_bias", _slice_conv) # [conv_dim]
if self.rmsnorm:
DMRegistry.convert(self.norm)
@@ -1151,7 +1181,8 @@ class _DynamicMambaMixer(DynamicModule):
"""Export the dynamic module to a torch.nn.Module."""
self.in_proj.export()
self.out_proj.export()
self.conv1d.export()
if hasattr(self, "conv1d"): # nemo:26.06 and earlier
self.conv1d.export()
if self.rmsnorm:
self.norm.export()
return super().export()
@@ -37,11 +37,17 @@ Hybrid pattern characters (from ``megatron.core.ssm.mamba_hybrid_layer_allocatio
import io
import sys
from typing import Any
from typing import TYPE_CHECKING, Any
import torch
from megatron.core.models.gpt.gpt_model import GPTModel
from megatron.core.models.mamba.mamba_model import MambaModel
try: # nemo:26.08+
from megatron.core.models.hybrid.hybrid_model import HybridModel
_HYBRID_MODEL_TYPES: tuple[type, ...] = (MambaModel, HybridModel)
except ImportError: # nemo:26.06 and earlier
_HYBRID_MODEL_TYPES = (MambaModel,)
from megatron.core.parallel_state import (
get_expert_tensor_and_model_parallel_group,
get_expert_tensor_parallel_rank,
@@ -56,6 +62,10 @@ from rich.table import Table
from modelopt.torch.opt.dynamic import DynamicModule
from modelopt.torch.utils import num2hrb, print_rank_0
if TYPE_CHECKING:
from megatron.core.models.gpt.gpt_model import GPTModel
from megatron.core.models.hybrid.hybrid_model import HybridModel # noqa: TC004
__all__ = [
"mcore_memory_footprint_mb",
"mcore_param_count",
@@ -558,8 +568,8 @@ def mcore_param_count(
return total, active
def mcore_param_count_live(model: GPTModel | MambaModel) -> int:
"""Count parameters in a live MCore GPTModel or MambaModel (reduced across TP, EP, ETP, and PP ranks)."""
def mcore_param_count_live(model: "GPTModel | MambaModel | HybridModel") -> int:
"""Count parameters in a live MCore LLM model (reduced across TP, EP, ETP, and PP ranks)."""
if isinstance(model, DynamicModule):
raise RuntimeError(
"mcore_param_count_live does not support DynamicModule. "
@@ -718,7 +728,7 @@ def mcore_memory_footprint_mb(
def print_mcore_model_stats(
model: "GPTModel | MambaModel",
model: "GPTModel | MambaModel | HybridModel",
label: str = "Model",
seq_length: int = 4096,
batch_size: int = 1,
@@ -727,7 +737,7 @@ def print_mcore_model_stats(
"""Print total params, active params, and memory footprint for an MCore model.
Args:
model: GPTModel or MambaModel to print stats for.
model: MCore LLM model to print stats for.
label: Label prefix for the output line (e.g. ``"Original"``, ``"Pruned"``).
seq_length: Sequence length for KV-cache / Mamba-state memory estimate.
batch_size: Batch size for KV-cache / Mamba-state memory estimate.
@@ -735,7 +745,7 @@ def print_mcore_model_stats(
"""
hybrid_layer_pattern: str | None = None
config_overrides: dict = {}
if isinstance(model, MambaModel):
if isinstance(model, _HYBRID_MODEL_TYPES):
hybrid_key = (
"hybrid_override_pattern"
if hasattr(model, "hybrid_override_pattern")
@@ -744,8 +754,8 @@ def print_mcore_model_stats(
hybrid_layer_pattern = getattr(model, hybrid_key)
# mamba_num_heads may not be stored in config when derived from model architecture;
# fall back to reading it from the actual layer.
if getattr(model.config, "mamba_num_heads", None) is None:
for layer in model.decoder.layers:
if getattr(model.config, "mamba_num_heads", None) is None: # type: ignore[attr-defined]
for layer in model.decoder.layers: # type: ignore[attr-defined]
if hasattr(layer, "mixer") and hasattr(layer.mixer, "nheads"):
config_overrides["mamba_num_heads"] = layer.mixer.nheads
break
@@ -168,7 +168,7 @@ def drop_mcore_language_model_layers(model: nn.Module, *, layers_to_drop: list[i
assert isinstance(model, supported_model_types), (
f"Model should have one of {supported_model_types} submodule, got {model}"
)
print_rank_0(f"Dropping decoder layers {layers_to_drop} from model.")
print_rank_0(f"Dropping decoder layers {layers_to_drop} (1-indexed) from model.")
# get the number of layers remaining in each pp rank
layers_remaining_per_pp = torch.zeros(
+8 -7
View File
@@ -763,7 +763,7 @@ def get_dataset_dataloader(
batch_size: Batch size of the returned dataloader.
num_samples: Number of samples from the dataset (interpreted as number of *output
rows* in both ``pack=False`` and ``pack=True`` modes — in packed mode the
loader oversamples raw text 4x to ensure enough docs to fill all rows).
loader oversamples raw text 16x to ensure enough docs to fill all rows).
max_sample_length: Maximum length of a sample (or per-row length under ``pack=True``).
device: Target device for the returned dataloader.
include_labels: Whether to include labels in the dataloader (ignored when
@@ -837,9 +837,9 @@ def get_dataset_dataloader(
# Sample count semantics:
# - pack=False: gather exactly `num_sample` raw docs per source, one per output row.
# - pack=True: oversample 8x per source to ensure enough raw docs to fill all rows,
# - pack=True: oversample 16x per source to ensure enough raw docs to fill all rows,
# since each row greedily packs multiple docs.
sample_multiplier = 8 if pack else 1
sample_multiplier = 16 if pack else 1
all_samples = []
for ds_name, num_sample in zip(dataset_name, num_samples):
samples = get_dataset_samples(
@@ -862,10 +862,11 @@ def get_dataset_dataloader(
)
if input_ids.shape[0] < total_rows:
warn_rank_0(
f"pack=True produced {input_ids.shape[0]} rows out of {total_rows} "
f"requested — raw text exhausted before filling all rows (8x oversample "
f"of num_samples was insufficient). Increase `num_samples` or shorten "
f"`max_sample_length`."
f"pack=True produced {input_ids.shape[0]} rows out of {total_rows} requested — "
f"raw text exhausted before filling all rows ({sample_multiplier}x oversample of "
"num_samples was insufficient). Shorten `max_sample_length` or supply longer, "
"more token-dense samples; increasing `num_samples` only helps if the source can "
"provide additional useful content."
)
if device:
input_ids = input_ids.to(device)
+7 -1
View File
@@ -15,6 +15,7 @@
"""Handles suppressing import errors for third-party modules that may or may not be available."""
import sys
from contextlib import contextmanager
from .logging import warn_rank_0
@@ -32,7 +33,12 @@ def import_plugin(plugin_name, msg_if_missing=None, verbose=True, success_msg=No
warn_rank_0(msg_if_missing)
except Exception as e:
if verbose:
# Capture the ``with import_plugin(...)`` call site so warnings point at the plugin
# that actually failed rather than at this helper. When ``__enter__`` runs the
# generator body, the stack is [0]=here, [1]=contextlib.__enter__, [2]=the caller.
caller = sys._getframe(2)
warn_rank_0(
f"Failed to import modelopt {plugin_name} plugin due to: {e!r}. "
f"Failed to import modelopt {plugin_name} plugin "
f"(from {caller.f_code.co_filename}:{caller.f_lineno}) due to: {e!r}. "
"You may ignore this warning if you do not need this plugin."
)
+23 -7
View File
@@ -32,9 +32,20 @@ from megatron.core.transformer.module import MegatronModule
from megatron.core.utils import unwrap_model
from transformers import AutoTokenizer
from modelopt.torch.nas.plugins.megatron import get_te_mamba_stack_spec
from modelopt.torch.nas.plugins.megatron import get_te_hybrid_stack_spec, get_te_mamba_stack_spec
from modelopt.torch.utils import print_rank_0
try: # nemo:26.08+
from megatron.bridge.models.hybrid.hybrid_provider import HybridModelProvider
from megatron.core.models.hybrid.hybrid_model import HybridModel
HAS_HYBRID = True
except ImportError:
HAS_HYBRID = False
HybridModelProvider = None
HybridModel = None
__all__ = ["load_mbridge_model_from_hf", "load_modelopt_megatron_checkpoint"]
@@ -48,9 +59,9 @@ def load_mbridge_model_from_hf(
load_weights: bool = True,
) -> tuple[
AutoBridge,
GPTModelProvider | MambaModelProvider,
GPTModelProvider | MambaModelProvider | HybridModelProvider,
list[MegatronModule],
GPTModel | MambaModel,
MegatronModule,
AutoTokenizer,
]:
"""Load a Megatron-Bridge model from HF.
@@ -88,8 +99,12 @@ def load_mbridge_model_from_hf(
# provider (the bridge's native, possibly custom/hybrid spec reads it at build time) rather than
# replacing the whole layer spec -- overwriting it would drop custom layers (e.g. Qwen3.5's
# GatedDeltaNet + gated-attention or Gemma3's custom spec).
if isinstance(provider, MambaModelProvider):
if HAS_HYBRID and isinstance(provider, (HybridModelProvider)):
provider.hybrid_stack_spec = get_te_hybrid_stack_spec(moe_grouped_gemm=moe_grouped_gemm)
provider.moe_grouped_gemm = moe_grouped_gemm
elif isinstance(provider, (MambaModelProvider)): # Deprecated in favor of HybridModelProvider
provider.mamba_stack_spec = get_te_mamba_stack_spec(moe_grouped_gemm=moe_grouped_gemm)
provider.moe_grouped_gemm = moe_grouped_gemm
elif (provider.num_moe_experts or 0) > 0:
provider.moe_grouped_gemm = moe_grouped_gemm
provider.finalize()
@@ -100,10 +115,11 @@ def load_mbridge_model_from_hf(
assert len(model) == 1
unwrapped_model = unwrap_model(model[0])
# VLMs (e.g. Qwen3-VL) wrap the language model as ``.language_model``; the pruning target is the
# inner GPTModel/MambaModel, but we still return the full wrapper so callers can save the VLM.
# inner GPTModel/MambaModel/HybridModel, but we still return the full wrapper so callers can save the VLM.
language_model = getattr(unwrapped_model, "language_model", unwrapped_model)
assert isinstance(language_model, (GPTModel, MambaModel)), (
f"Expected a GPTModel/MambaModel (optionally wrapped as .language_model), "
model_types = (GPTModel, MambaModel, HybridModel) if HAS_HYBRID else (GPTModel, MambaModel)
assert isinstance(language_model, model_types), (
f"Expected a GPTModel/MambaModel/HybridModel (optionally wrapped as `.language_model`), "
f"got {type(unwrapped_model)}"
)
+47 -41
View File
@@ -36,7 +36,7 @@ from megatron.core.transformer.module import MegatronModule
from megatron.core.transformer.transformer_config import MLATransformerConfig, TransformerConfig
from modelopt.torch.export.unified_export_megatron import import_mcore_gpt_from_hf
from modelopt.torch.nas.plugins.megatron import get_te_mamba_stack_spec
from modelopt.torch.nas.plugins.megatron import get_te_hybrid_stack_spec, get_te_mamba_stack_spec
try:
from megatron.core.extensions.transformer_engine import TENorm
@@ -57,6 +57,16 @@ except ImportError as e:
warn(f"Mamba not installed: {e}")
HAS_MAMBA = False
try: # nemo:26.08+ (MambaModel is now a deprecated HybridModel subclass)
from megatron.core.models.hybrid.hybrid_model import HybridModel
from megatron.core.post_training.modelopt.hybrid.model_specs import (
get_hybrid_stack_modelopt_spec,
)
HAS_HYBRID = True
except ImportError:
HAS_HYBRID = False
try:
import apex # noqa: F401
@@ -360,7 +370,7 @@ def get_mcore_mamba_hybrid_model(
num_layers: int = 3,
num_layers_in_first_pipeline_stage: int | None = None,
num_layers_in_last_pipeline_stage: int | None = None,
hybrid_override_pattern: str | None = None,
hybrid_layer_pattern: str | None = None,
hidden_size: int = 64,
num_attention_heads: int = 8,
num_query_groups: int | None = None,
@@ -386,7 +396,7 @@ def get_mcore_mamba_hybrid_model(
"""Builds a Mamba model with hybrid layer allocation (Mamba, MoE, Attention, MLP blocks).
Notable Args:
hybrid_override_pattern: The hybrid layer pattern to override with.
hybrid_layer_pattern: The hybrid layer pattern to use.
If None, a default pattern will be generated.
skip_moe: Whether to skip MoE blocks in default hybrid pattern.
"""
@@ -422,49 +432,45 @@ def get_mcore_mamba_hybrid_model(
**config_kwargs,
)
# TODO: hybrid_override_pattern is deprecated in MCore 0.17+, use hybrid_layer_pattern instead
if hybrid_override_pattern is None:
if hybrid_layer_pattern is None:
# Generate pattern by repeating base_pattern and trimming to match num_layers
# E.g. for num_layers=3, return "MEM" (Mamba -> MoE -> Mamba)
# E.g. for num_layers=6, return "MEM*M-" (Mamba -> MoE -> Attention -> MoE -> MLP)
base_pattern = "M*M-" if skip_moe else "MEM*M-"
hybrid_override_pattern = (base_pattern * num_layers)[:num_layers]
hybrid_layer_pattern = (base_pattern * num_layers)[:num_layers]
# TODO: enable this when MCore 0.17+ is released (has fall-back so without this is still fine for sometime)
# Add | symbols for Pipeline parallelism (supported from MCore 0.17+, auto-added if not provided)
# E.g. MEM* with PP2 becomes ME|M* and MEM*M-ME with PP2 becomes MEM*|M-ME
# if pipeline_model_parallel_size > 1:
# if "|" not in hybrid_override_pattern:
# assert (
# num_layers_in_first_pipeline_stage is None
# and num_layers_in_last_pipeline_stage is None
# ), "hybrid_override_pattern with `|` must be provided for uneven PP"
# hybrid_override_pattern = "|".join(
# textwrap.wrap(
# hybrid_override_pattern,
# width=num_layers // pipeline_model_parallel_size,
# break_long_words=True,
# break_on_hyphens=False,
# )
# )
# assert hybrid_override_pattern.count("|") == pipeline_model_parallel_size - 1
assert len(hybrid_override_pattern.replace("|", "")) == num_layers
print(f"Using `{hybrid_override_pattern=}` for building MambaModel")
# NOTE: We intentionally keep hybrid_layer_pattern pipe-free even under PP>1 (MCore warns and
# runtime-slices). This matches how bridge-loaded models are pruned; ModelOpt's depth-pruning
# slicer is not `|`-aware, so pipe stage separators would break pattern slicing -- reject them.
assert "|" not in hybrid_layer_pattern, "Pipeline separators (`|`) are not supported"
assert len(hybrid_layer_pattern) == num_layers
if transformer_impl == "transformer_engine":
mamba_spec = get_te_mamba_stack_spec(moe_grouped_gemm=moe_grouped_gemm)
# nemo:26.08+ uses HybridModel + hybrid_layer_pattern; older uses deprecated MambaModel.
common_kwargs = {
"config": config,
"vocab_size": vocab_size,
"max_sequence_length": max_sequence_length,
"pre_process": is_pipeline_first_stage(),
"post_process": is_pipeline_last_stage(),
"share_embeddings_and_output_weights": False,
"position_embedding_type": "none",
}
if HAS_HYBRID:
spec = (
get_te_hybrid_stack_spec(moe_grouped_gemm)
if transformer_impl == "transformer_engine"
else get_hybrid_stack_modelopt_spec(remap_te_layernorm=True)
)
model = HybridModel(
hybrid_stack_spec=spec, hybrid_layer_pattern=hybrid_layer_pattern, **common_kwargs
)
else:
mamba_spec = get_mamba_stack_modelopt_spec(remap_te_layernorm=True)
model = MambaModel(
config=config,
mamba_stack_spec=mamba_spec,
vocab_size=vocab_size,
max_sequence_length=max_sequence_length,
hybrid_override_pattern=hybrid_override_pattern,
pre_process=is_pipeline_first_stage(),
post_process=is_pipeline_last_stage(),
share_embeddings_and_output_weights=False,
position_embedding_type="none",
)
spec = (
get_te_mamba_stack_spec(moe_grouped_gemm)
if transformer_impl == "transformer_engine"
else get_mamba_stack_modelopt_spec(remap_te_layernorm=True)
)
model = MambaModel(
mamba_stack_spec=spec, hybrid_override_pattern=hybrid_layer_pattern, **common_kwargs
)
return model.to(torch.bfloat16) if bf16 else model
+9 -2
View File
@@ -43,6 +43,13 @@ from modelopt.torch.opt.plugins.mcore_dist_checkpointing import (
)
from modelopt.torch.utils import to_empty_if_meta_device
try: # nemo:26.08+: MambaModel is a HybridModel subclass; toy models build HybridModel directly.
from megatron.core.models.hybrid.hybrid_model import HybridModel
_MAMBA_HYBRID_TYPES: tuple[type, ...] = (MambaModel, HybridModel)
except ImportError:
_MAMBA_HYBRID_TYPES = (MambaModel,)
@torch.no_grad()
def run_mcore_inference(
@@ -129,8 +136,8 @@ def get_forward(model, batch_size=2):
input_ids, labels, position_ids, attention_mask, loss_mask = get_batch(model, batch_size)
def forward(model):
# MambaModel doesn't accept loss_mask argument
if isinstance(model, MambaModel):
# Mamba/Hybrid forward doesn't accept loss_mask argument
if isinstance(model, _MAMBA_HYBRID_TYPES):
return model.forward(
input_ids=input_ids,
position_ids=position_ids,
@@ -63,7 +63,7 @@ def _build_mcore_nemotron_h(config, size, initialize=True):
pipeline_model_parallel_size=1,
initialize_megatron=initialize,
num_layers=config.num_hidden_layers,
hybrid_override_pattern=config.hybrid_override_pattern,
hybrid_layer_pattern=config.hybrid_override_pattern,
hidden_size=config.hidden_size,
num_attention_heads=config.num_attention_heads,
num_query_groups=config.num_key_value_heads,
@@ -51,7 +51,7 @@ def _test_mamba_search_space(rank, size):
mamba_head_dim_divisor = 4
num_layers = size
hybrid_override_pattern = "M" * size # all layers are Mamba layers
hybrid_layer_pattern = "M" * size # all layers are Mamba layers
hidden_size = channel_divisor * 4
mamba_state_dim = channel_divisor
mamba_head_dim = mamba_head_dim_divisor * 2
@@ -65,7 +65,7 @@ def _test_mamba_search_space(rank, size):
pipeline_model_parallel_size=size,
initialize_megatron=True,
num_layers=num_layers,
hybrid_override_pattern=hybrid_override_pattern,
hybrid_layer_pattern=hybrid_layer_pattern,
hidden_size=hidden_size,
mamba_state_dim=mamba_state_dim,
mamba_head_dim=mamba_head_dim,
@@ -100,7 +100,8 @@ def _test_mamba_search_space(rank, size):
assert isinstance(layer.mixer, _DynamicMambaMixer)
assert isinstance(layer.mixer.in_proj, _DynamicTELayerNormColumnParallelLinear)
assert isinstance(layer.mixer.out_proj, _DynamicTERowParallelLinear)
assert isinstance(layer.mixer.conv1d, _DynamicConvNd)
if hasattr(layer.mixer, "conv1d"): # nemo:26.06 and earlier
assert isinstance(layer.mixer.conv1d, _DynamicConvNd)
if layer.mixer.rmsnorm:
assert isinstance(layer.mixer.norm, _DynamicExtendedRMSNorm)
if is_pipeline_last_stage():
@@ -523,7 +523,7 @@ def _test_formula_matches_mamba_model(rank, size, parallelism):
mamba_head_dim=mamba_head_dim,
mamba_num_heads=mamba_num_heads,
vocab_size=128,
hybrid_override_pattern=pattern,
hybrid_layer_pattern=pattern,
moe_grouped_gemm=False,
num_moe_experts=4,
moe_ffn_hidden_size=64,
@@ -53,7 +53,7 @@ def _test_mcore_mamba_parameter_sorting(rank, size):
channel_divisor = 64
num_layers = size
hybrid_override_pattern = "M" * size
hybrid_layer_pattern = "M" * size
hidden_size = channel_divisor * 4
mamba_state_dim = channel_divisor
mamba_head_dim = 16
@@ -67,7 +67,7 @@ def _test_mcore_mamba_parameter_sorting(rank, size):
pipeline_model_parallel_size=size,
initialize_megatron=True,
num_layers=num_layers,
hybrid_override_pattern=hybrid_override_pattern,
hybrid_layer_pattern=hybrid_layer_pattern,
hidden_size=hidden_size,
mamba_state_dim=mamba_state_dim,
mamba_head_dim=mamba_head_dim,
@@ -210,7 +210,10 @@ def _test_mcore_mamba_hybrid_pruning(rank, size, ckpt_dir):
assert mixer.headdim == pruned_mamba_head_dim
assert mixer.d_inner == pruned_mamba_num_heads * pruned_mamba_head_dim
assert mixer.out_proj.out_features == pruned_hidden_size
assert mixer.conv1d.in_channels == mixer.conv1d.out_channels == mixer.d_inner + bc
if hasattr(mixer, "conv1d"): # nemo:26.06 and earlier
assert mixer.conv1d.in_channels == mixer.conv1d.out_channels == mixer.d_inner + bc
else: # nemo:26.08+
assert mixer.conv1d_weight.shape[0] == mixer.conv1d_bias.shape[0] == mixer.d_inner + bc
# Assert model.config is updated for correct save/restoring
assert model.config.ffn_hidden_size == pruned_ffn_hidden_size
@@ -239,7 +242,7 @@ _NAS_CHANNEL_DIVISOR = 4
_NAS_BATCH_SIZE = 2
_NAS_MODEL_KWARGS = {
"num_layers": 4,
"hybrid_override_pattern": "ME*-",
"hybrid_layer_pattern": "ME*-",
"hidden_size": 16,
"ffn_hidden_size": 32,
"num_attention_heads": 16,
@@ -267,7 +267,7 @@ def _gpt_model_provider(
transformer_impl="local",
# Hybrid mamba MOE parameters
is_hybrid=False,
hybrid_override_pattern=None,
hybrid_layer_pattern=None,
mamba_head_dim=16,
):
device_ctx = torch.device("meta") if meta_device else nullcontext()
@@ -275,7 +275,7 @@ def _gpt_model_provider(
with device_ctx:
if is_hybrid:
# Derive num_layers from pattern length, default to 4
num_layers = len(hybrid_override_pattern) if hybrid_override_pattern else 4
num_layers = len(hybrid_layer_pattern) if hybrid_layer_pattern else 4
model = get_mcore_mamba_hybrid_model(
tensor_model_parallel_size=tp_size,
num_layers=num_layers,
@@ -283,7 +283,7 @@ def _gpt_model_provider(
vocab_size=vocab_size,
num_attention_heads=8,
ffn_hidden_size=None,
hybrid_override_pattern=hybrid_override_pattern,
hybrid_layer_pattern=hybrid_layer_pattern,
mamba_head_dim=mamba_head_dim,
mamba_num_groups=tp_size, # Must be divisible by tp_size
num_moe_experts=num_moe_experts,
@@ -333,7 +333,7 @@ def _test_sharded_state_dict(
transformer_impl = model_config.get("transformer_impl", "local")
# Hybrid mamba MOE parameters
is_hybrid = model_config.get("is_hybrid", False)
hybrid_override_pattern = model_config.get("hybrid_override_pattern", None)
hybrid_layer_pattern = model_config.get("hybrid_layer_pattern", None)
initialize_for_megatron(
tensor_model_parallel_size=tp_size,
@@ -352,7 +352,7 @@ def _test_sharded_state_dict(
etp_size=etp_size,
transformer_impl=transformer_impl,
is_hybrid=is_hybrid,
hybrid_override_pattern=hybrid_override_pattern,
hybrid_layer_pattern=hybrid_layer_pattern,
)
model_test = _gpt_model_provider(
tp_size,
@@ -365,7 +365,7 @@ def _test_sharded_state_dict(
etp_size=etp_size,
transformer_impl=transformer_impl,
is_hybrid=is_hybrid,
hybrid_override_pattern=hybrid_override_pattern,
hybrid_layer_pattern=hybrid_layer_pattern,
)
forward = get_forward(model_ref)
@@ -534,7 +534,7 @@ def test_homogeneous_sharded_state_dict_hybrid(dist_workers, tmp_path, config):
pytest.skip("Test needs to be fixed for more than 4 GPUs")
model_config = {
"is_hybrid": True,
"hybrid_override_pattern": "MEM*E", # 5 layers: Mamba → MoE → Mamba → Attention → MoE
"hybrid_layer_pattern": "MEM*E", # 5 layers: Mamba → MoE → Mamba → Attention → MoE
"num_moe_experts": 8,
"tp_size": num_gpus,
"ep_size": 1,
+1 -1
View File
@@ -699,7 +699,7 @@ def test_multi_source_pack_shuffles_to_avoid_dominance(monkeypatch, tiny_tokeniz
"""With ``pack=True`` and 2+ sources, samples are shuffled so a long-doc source
can't silently exhaust the row budget and drop the other sources.
Without shuffle, source A's 8x-oversampled docs would all come first in
Without shuffle, source A's 16x-oversampled docs would all come first in
``all_samples`` and (with sufficient row consumption per doc) fill every row.
With the deterministic shuffle, both sources appear within the first
``total_rows`` worth of consumed samples.
@@ -16,7 +16,7 @@
# limitations under the License.
# Megatron-Bridge HF -> Megatron checkpoint import.
# Assumes nvcr.io/nvidia/nemo:26.02+ container (megatron-bridge preinstalled at /opt/Megatron-Bridge).
# Assumes nvcr.io/nvidia/nemo:26.04+ container (megatron-bridge preinstalled at /opt/Megatron-Bridge).
#
# Required env: HF_MODEL_ID (e.g. nvidia/NVIDIA-Nemotron-3-Super-120B-A12B-BF16)
# Optional env: