Skip to content
Open
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
9 changes: 6 additions & 3 deletions src/mcore_bridge/model/modules/mtp_layer.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@
from megatron.core import InferenceParams, parallel_state, tensor_parallel
from megatron.core.enums import Fp8Recipe
from megatron.core.extensions.transformer_engine import te_checkpoint
from megatron.core.fp4_utils import get_fp4_context
from megatron.core.fp8_utils import get_fp8_context
from megatron.core.packed_seq_params import PackedSeqParams
from megatron.core.tensor_parallel.mappings import (gather_from_sequence_parallel_region,
Expand Down Expand Up @@ -343,16 +344,18 @@ def _proj_and_transformer_layer(
else:
rng_context = nullcontext()

# Unlike transformer_block.py which needs to support mixed-precision in
# different layers,currently MTP only use global fp8 context.
# Same quant context as TransformerBlock. Leaving FP4 as a null context
# trains the MTP loss in the default dtype while the decoder is FP4.
if self.config.fp8:
fp8_context = get_fp8_context(self.config)
transformer_layer_fp8_context = get_fp8_context(self.config)
elif self.config.fp4:
fp8_context = get_fp4_context(self.config)
transformer_layer_fp8_context = get_fp4_context(self.config)
else:
fp8_context = nullcontext()
transformer_layer_fp8_context = nullcontext()

# TODO: currently ignoring FP4 in MTP layers because we need more numerical validation
with rng_context:
with fp8_context:
hidden_states = self._concat_embeddings(hidden_states, decoder_input)
Expand Down