Skip to content
Draft
8 changes: 5 additions & 3 deletions hybrid_builders.py
Original file line number Diff line number Diff line change
@@ -1,12 +1,12 @@
# Copyright (c) 2025-2026, NVIDIA CORPORATION. All rights reserved.

from model_provider import count_parameters_in_layer
from megatron.core.models.hybrid.hybrid_layer_specs import hybrid_inference_stack_spec
from megatron.core.models.hybrid.hybrid_model import HybridModel
from megatron.core.transformer import TransformerConfig
from megatron.core.transformer.spec_utils import import_module
from megatron.core.transformer.spec_utils import ModuleSpec, import_module
from megatron.training import print_rank_0
from megatron.training.arguments import core_transformer_config_from_args
from megatron.core.models.hybrid.hybrid_layer_specs import hybrid_inference_stack_spec
from model_provider import count_parameters_in_layer


def hybrid_builder(args, pre_process, post_process, vp_stage=None, config=None, pg_collection=None):
Expand All @@ -21,6 +21,8 @@ def hybrid_builder(args, pre_process, post_process, vp_stage=None, config=None,
), "inference_fuse_tp_communication is not supported for HybridModel"
elif args.spec is not None:
hybrid_stack_spec = import_module(args.spec)
if callable(hybrid_stack_spec) and not isinstance(hybrid_stack_spec, ModuleSpec):
hybrid_stack_spec = hybrid_stack_spec(config)
else:
raise ValueError("You must provide a valid hybrid layer spec via --spec")

Expand Down
27 changes: 27 additions & 0 deletions megatron/core/fp8_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -871,6 +871,29 @@ def get_fp8_context(config: TransformerConfig, layer_no: int = -1, is_init: bool

return fp8_context

def get_fp8_disabled_context(config: TransformerConfig, is_init: bool = False):
"""Return a context manager that disables TE quantization.

Use this around submodule construction or execution that must stay in a higher
precision while its enclosing module uses an FP8 or FP4 context.

Args:
config: Transformer configuration that controls quantization.
is_init: Whether to disable the parameter-initialization context instead of
the forward autocast context.

Returns:
A disabled TE quantization context when quantization is active, otherwise a
no-op context.
"""
if is_init:
if not (config.fp8_param or config.fp4_param):
return nullcontext()
return transformer_engine.pytorch.fp8_model_init(enabled=False)
if not (config.fp8 or config.fp4):
return nullcontext()
return transformer_engine.pytorch.fp8_autocast(enabled=False)

else:

def get_fp8_recipe(config: TransformerConfig):
Expand All @@ -881,6 +904,10 @@ def get_fp8_context(config: TransformerConfig, layer_no: int = -1, is_init: bool
"""Returns dummy fp8 context manager since TE is not available."""
return nullcontext()

def get_fp8_disabled_context(config: TransformerConfig, is_init: bool = False):
"""Return a no-op context manager since TE is not available."""
return nullcontext()


if HAVE_TE:
from transformer_engine.pytorch.fp8 import FP8GlobalStateManager
Expand Down
93 changes: 90 additions & 3 deletions megatron/core/fusions/fused_bias_dropout.py
Original file line number Diff line number Diff line change
@@ -1,10 +1,13 @@
# Copyright (c) 2023, NVIDIA CORPORATION. All rights reserved.
from typing import Optional, Tuple
# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
from typing import TYPE_CHECKING, Optional, Tuple

import torch

from megatron.core.jit import jit_fuser

if TYPE_CHECKING:
from megatron.core.tensor_parallel.random import CheckpointWithoutOutputManager

# pylint: disable=missing-function-docstring


Expand Down Expand Up @@ -80,7 +83,26 @@ def bias_dropout_add_fused_inference(
return _bias_dropout_add_func(x_with_bias, residual, prob, False)


def get_bias_dropout_add(training, fused):
def get_bias_dropout_add(
training, fused, mhc_recompute_manager: Optional['CheckpointWithoutOutputManager'] = None
):
"""
Get the bias-dropout-add function.

Args:
training: Whether in training mode.
fused: Whether to use fused implementation.
mhc_recompute_manager: Optional CheckpointWithoutOutputManager for checkpoint management.
When provided, the returned function will wrap the BDA operation with
CheckpointWithoutOutput for memory-efficient recomputation.

Returns:
A callable that performs bias-dropout-add operation.
"""
if mhc_recompute_manager is not None:
# Return a checkpointed version that handles tuple unpacking internally
return _get_checkpointed_bda(training, fused, mhc_recompute_manager)

if fused:
# jit scripting for a nn.module (with dropout) is not
# triggering the fusion kernel. For now, we use two
Expand All @@ -92,3 +114,68 @@ def get_bias_dropout_add(training, fused):
return bias_dropout_add_fused_inference
else:
return bias_dropout_add_unfused(training)


def _get_checkpointed_bda(training, fused, mhc_recompute_manager: 'CheckpointWithoutOutputManager'):
"""
Create a checkpointed bias-dropout-add function.

This function handles:
1. Tuple unpacking for x_with_bias (required because save_for_backward can't save tuples)
2. Non-tensor arguments like dropout probability (handled by CheckpointWithoutOutput)
3. Auto-registration to the CheckpointWithoutOutputManager

Args:
training: Whether in training mode.
fused: Whether to use fused implementation.
mhc_recompute_manager: CheckpointWithoutOutputManager for checkpoint management.

Returns:
A callable that performs checkpointed bias-dropout-add operation.
"""
from megatron.core.tensor_parallel.random import CheckpointWithoutOutput

# Get the underlying BDA function
if fused:
if training:
bda_func = bias_dropout_add_fused_train
else:
bda_func = bias_dropout_add_fused_inference
else:
bda_func = bias_dropout_add_unfused(training)

def _checkpointed_bda(x_with_bias, residual, prob):
"""
Checkpointed BDA that handles tuple unpacking internally.

Args:
x_with_bias: Either a tuple (x, bias) or a single tensor x.
residual: Residual tensor.
prob: Dropout probability.

Returns:
Output tensor after bias-dropout-add.
"""
# Create checkpoint with manager
ckpt = CheckpointWithoutOutput(ckpt_manager=mhc_recompute_manager)

# Handle case where x_with_bias might be a single tensor (e.g., from IdentityOp)
if isinstance(x_with_bias, tuple):
x, bias = x_with_bias
else:
x = x_with_bias
bias = None

# Wrapper function that re-packs the tuple for the actual BDA function
def _bda_wrapper(output, bias, res, dropout):
return bda_func((output, bias), res, dropout)

# Call checkpoint with unpacked arguments
result = ckpt.checkpoint(_bda_wrapper, x, bias, residual, prob)

# No-op when manager is set - manager handles all discarding uniformly
ckpt.discard_output_and_register_recompute(result)

return result

return _checkpointed_bda
Loading