Repository navigation
Conversation
The MTP forward wrapped the block in get_fp8_context only when config.fp8 was set. With FP4 enabled the context stayed null, so the MTP loss ran in the default dtype while the decoder ran in FP4. MTP now uses get_fp4_context for both the embedding concatenation and the inner transformer layer when config.fp4 is set, matching TransformerBlock. FP8 and the unquantized path are unchanged. flake8 is clean on mtp_layer.py. This path needs Megatron's fp4_utils, which is already imported by transformer_block.py.
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
The MTP forward wrapped the block in
get_fp8_contextonly whenconfig.fp8was set. With FP4 enabled the context stayed null, so the MTP loss ran in the default dtype while the decoder ran in FP4.MTP now uses
get_fp4_contextfor both the embedding concatenation and the inner transformer layer whenconfig.fp4is set, matchingTransformerBlock. FP8 and the unquantized path are unchanged.flake8is clean onmtp_layer.py. This path needs Megatron'sfp4_utils, which is already imported bytransformer_block.py.