diff --git a/tests/models/qwen4_exp/test_ple.py b/tests/models/qwen4_exp/test_ple.py index 35c24176f8e5..d7747ebfbdb5 100644 --- a/tests/models/qwen4_exp/test_ple.py +++ b/tests/models/qwen4_exp/test_ple.py @@ -14,6 +14,9 @@ import vllm.model_executor.layers.vocab_parallel_embedding as embedding_module import vllm.model_executor.parameter as parameter_module from vllm.model_executor.layers.quantization.fp8 import Fp8Config +from vllm.model_executor.layers.quantization.modelopt import ( + ModelOptMixedPrecisionConfig, +) from vllm.models.qwen4_exp.common.ple import ( PLEShardOverlap, compute_ple_shard_overlap, @@ -288,6 +291,37 @@ def test_ple_fp8_embedding_respects_checkpoint_shard_exclusions() -> None: assert _get_ple_embedding_quant_method(quant_config, prefix) is None +def test_ple_fp8_embedding_supports_mixed_precision_config() -> None: + prefix = "model.language_model.layers.1.ple.ple_embedding.ngram_embedding" + quant_config = ModelOptMixedPrecisionConfig.from_config( + { + "quantization": { + "quant_algo": "MIXED_PRECISION", + "exclude_modules": [], + "group_size": 16, + "quantized_layers": { + prefix: {"quant_algo": "FP8"}, + "model.language_model.layers.2.moe.gate_proj": { + "quant_algo": "NVFP4" + }, + }, + } + } + ) + + assert isinstance( + _get_ple_embedding_quant_method(quant_config, prefix), + Qwen4ExpPLEFp8EmbeddingMethod, + ) + assert ( + _get_ple_embedding_quant_method( + quant_config, + "model.language_model.layers.2.moe.gate_proj", + ) + is None + ) + + def test_dilated_ple_spec_state_rolls_back_before_next_forward() -> None: conv_state_len = 6 dilation = 2 diff --git a/vllm/models/qwen4_exp/nvidia/ple_layer.py b/vllm/models/qwen4_exp/nvidia/ple_layer.py index c0d909e1c384..9d7091d7a2b3 100644 --- a/vllm/models/qwen4_exp/nvidia/ple_layer.py +++ b/vllm/models/qwen4_exp/nvidia/ple_layer.py @@ -22,6 +22,9 @@ QuantizeMethodBase, ) from vllm.model_executor.layers.quantization.fp8 import Fp8Config +from vllm.model_executor.layers.quantization.modelopt import ( + ModelOptMixedPrecisionConfig, +) from vllm.model_executor.layers.quantization.utils.fp8_utils import ( create_fp8_scale_parameter, create_fp8_weight_parameter, @@ -133,6 +136,11 @@ def _get_ple_embedding_quant_method( ) -> QuantizeMethodBase | None: """Select global-scale FP8 only for quantized PLE checkpoint shards.""" + if isinstance(quant_config, ModelOptMixedPrecisionConfig): + if quant_config._resolve_quant_algo(prefix) == "FP8": + return Qwen4ExpPLEFp8EmbeddingMethod() + return None + if not isinstance(quant_config, Fp8Config): return None if not quant_config.is_checkpoint_fp8_serialized: