Skip to content

fix: run MTP layers in the FP4 context when FP4 is enabled - #224

Open
shiaho777 wants to merge 1 commit into
modelscope:mainfrom
shiaho777:fix/mtp-fp4-context
Open

shiaho777 wants to merge 1 commit into
modelscope:mainfrom
shiaho777:fix/mtp-fp4-context

Conversation

@shiaho777

Copy link
Copy Markdown
Contributor

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.

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.
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant