fix(moe): support Llama4 NVFP4 router input weights on SM120 (#35504)
Co-authored-by: Po-Han Huang (NVIDIA) <53919306+nvpohanh@users.noreply.github.com>
This commit is contained in:
co-authored by
Po-Han Huang
parent
59a723ef1e
commit
56fee88e23
@@ -160,6 +160,36 @@ def _maybe_apply_routed_scaling_factor(
|
|||||||
return output
|
return output
|
||||||
|
|
||||||
|
|
||||||
|
def _prescale_router_weight_on_input(
|
||||||
|
dispatch_output: StandardDispatchOutput | FlashinferDispatchOutput,
|
||||||
|
runner_config: MoeRunnerConfig,
|
||||||
|
) -> StandardDispatchOutput | FlashinferDispatchOutput:
|
||||||
|
if not runner_config.apply_router_weight_on_input:
|
||||||
|
return dispatch_output
|
||||||
|
|
||||||
|
topk_output = dispatch_output.topk_output
|
||||||
|
topk_weights = topk_output.topk_weights
|
||||||
|
|
||||||
|
if dispatch_output.hidden_states_scale is not None:
|
||||||
|
raise NotImplementedError(
|
||||||
|
"apply_router_weight_on_input is not supported when activations are "
|
||||||
|
"quantized before dispatch (flashinfer_cutlass fp4 all-gather path)."
|
||||||
|
)
|
||||||
|
|
||||||
|
assert topk_weights.dim() == 2 and topk_weights.shape[-1] == 1, (
|
||||||
|
"apply_router_weight_on_input requires topk=1"
|
||||||
|
)
|
||||||
|
|
||||||
|
hidden_states = dispatch_output.hidden_states * topk_weights.to(
|
||||||
|
dispatch_output.hidden_states.dtype
|
||||||
|
)
|
||||||
|
unit_scales = torch.ones_like(topk_weights, dtype=torch.float32)
|
||||||
|
return dispatch_output._replace(
|
||||||
|
hidden_states=hidden_states,
|
||||||
|
topk_output=topk_output._replace(topk_weights=unit_scales),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def _prepare_input(
|
def _prepare_input(
|
||||||
dispatch_output,
|
dispatch_output,
|
||||||
quant_info: FlashInferCutlassMoeQuantInfo,
|
quant_info: FlashInferCutlassMoeQuantInfo,
|
||||||
@@ -197,8 +227,9 @@ def _run_flashinfer_cutlass(
|
|||||||
) -> torch.Tensor:
|
) -> torch.Tensor:
|
||||||
flashinfer_cutlass_fused_moe, _ = _flashinfer_cutlass_fused_moe()
|
flashinfer_cutlass_fused_moe, _ = _flashinfer_cutlass_fused_moe()
|
||||||
|
|
||||||
|
dispatch_output = _prescale_router_weight_on_input(dispatch_output, runner_config)
|
||||||
topk_output = dispatch_output.topk_output
|
topk_output = dispatch_output.topk_output
|
||||||
topk_weights = topk_output.topk_weights
|
topk_weights = topk_output.topk_weights.to(torch.float32)
|
||||||
topk_ids = topk_output.topk_ids
|
topk_ids = topk_output.topk_ids
|
||||||
x, x_sf, output_dtype, output_col = _prepare_input(
|
x, x_sf, output_dtype, output_col = _prepare_input(
|
||||||
dispatch_output, quant_info, runner_config
|
dispatch_output, quant_info, runner_config
|
||||||
@@ -267,9 +298,6 @@ def fused_experts_none_to_flashinfer_cutlass(
|
|||||||
assert isinstance(quant_info, FlashInferCutlassMoeQuantInfo), (
|
assert isinstance(quant_info, FlashInferCutlassMoeQuantInfo), (
|
||||||
f"Unexpected quant_info type for flashinfer_cutlass: {type(quant_info)}"
|
f"Unexpected quant_info type for flashinfer_cutlass: {type(quant_info)}"
|
||||||
)
|
)
|
||||||
assert not runner_config.apply_router_weight_on_input, (
|
|
||||||
"apply_router_weight_on_input is not supported for FlashInfer CUTLASS"
|
|
||||||
)
|
|
||||||
|
|
||||||
output = _run_flashinfer_cutlass(
|
output = _run_flashinfer_cutlass(
|
||||||
dispatch_output=dispatch_output,
|
dispatch_output=dispatch_output,
|
||||||
@@ -292,9 +320,6 @@ def fused_experts_flashinfer_to_flashinfer_cutlass(
|
|||||||
assert isinstance(quant_info, FlashInferCutlassMoeQuantInfo), (
|
assert isinstance(quant_info, FlashInferCutlassMoeQuantInfo), (
|
||||||
f"Unexpected quant_info type for flashinfer_cutlass: {type(quant_info)}"
|
f"Unexpected quant_info type for flashinfer_cutlass: {type(quant_info)}"
|
||||||
)
|
)
|
||||||
assert not runner_config.apply_router_weight_on_input, (
|
|
||||||
"apply_router_weight_on_input is not supported for FlashInfer CUTLASS"
|
|
||||||
)
|
|
||||||
|
|
||||||
output = _run_flashinfer_cutlass(
|
output = _run_flashinfer_cutlass(
|
||||||
dispatch_output=dispatch_output,
|
dispatch_output=dispatch_output,
|
||||||
|
|||||||
@@ -3053,7 +3053,6 @@ class ModelOptNvFp4FusedMoEMethod(FusedMoEMethodBase):
|
|||||||
assert activation in _SUPPORTED_ACT_STRS or (
|
assert activation in _SUPPORTED_ACT_STRS or (
|
||||||
activation == "situ" and moe_runner_backend.is_flashinfer_trtllm()
|
activation == "situ" and moe_runner_backend.is_flashinfer_trtllm()
|
||||||
), f"{activation=} is unsupported by {moe_runner_backend}"
|
), f"{activation=} is unsupported by {moe_runner_backend}"
|
||||||
moe_runner_config = self.moe_runner_config
|
|
||||||
|
|
||||||
if moe_runner_backend.is_flashinfer_megamoe():
|
if moe_runner_backend.is_flashinfer_megamoe():
|
||||||
from sglang.srt.layers.moe.flashinfer_megamoe import (
|
from sglang.srt.layers.moe.flashinfer_megamoe import (
|
||||||
@@ -3179,9 +3178,6 @@ class ModelOptNvFp4FusedMoEMethod(FusedMoEMethodBase):
|
|||||||
FlashInferCutlassMoeQuantInfo,
|
FlashInferCutlassMoeQuantInfo,
|
||||||
)
|
)
|
||||||
|
|
||||||
assert not moe_runner_config.apply_router_weight_on_input, (
|
|
||||||
"apply_router_weight_on_input is not supported for Flashinfer"
|
|
||||||
)
|
|
||||||
quant_info = FlashInferCutlassMoeQuantInfo(
|
quant_info = FlashInferCutlassMoeQuantInfo(
|
||||||
quant_type="fp4",
|
quant_type="fp4",
|
||||||
w13_weight=layer.w13_weight,
|
w13_weight=layer.w13_weight,
|
||||||
|
|||||||
Reference in New Issue
Block a user