Add explicit disable flag for FlashInfer allreduce fusion (#21446)
This commit is contained in:
@@ -518,6 +518,7 @@ class ServerArgs:
|
||||
moe_runner_backend: str = "auto"
|
||||
flashinfer_mxfp4_moe_precision: Literal["default", "bf16"] = "default"
|
||||
enable_flashinfer_allreduce_fusion: bool = False
|
||||
enforce_disable_flashinfer_allreduce_fusion: bool = False
|
||||
enable_aiter_allreduce_fusion: bool = False
|
||||
deepep_mode: Literal["auto", "normal", "low_latency"] = "auto"
|
||||
ep_num_redundant_experts: int = 0
|
||||
@@ -2082,6 +2083,14 @@ class ServerArgs:
|
||||
f"Auto-enabling FlashInfer AllReduce Fusion on SM90/SM10X for {model_arch}"
|
||||
)
|
||||
|
||||
# Apply enforce_disable_flashinfer_allreduce_fusion after all model-specific adjustments
|
||||
if self.enforce_disable_flashinfer_allreduce_fusion:
|
||||
self.enable_flashinfer_allreduce_fusion = False
|
||||
logger.info(
|
||||
"FlashInfer allreduce fusion is forcibly disabled "
|
||||
"via --enforce-disable-flashinfer-allreduce-fusion."
|
||||
)
|
||||
|
||||
def _handle_mamba_radix_cache(
|
||||
self,
|
||||
model_arch: str,
|
||||
@@ -4874,6 +4883,11 @@ class ServerArgs:
|
||||
action="store_true",
|
||||
help="Enable FlashInfer allreduce fusion with Residual RMSNorm.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--enforce-disable-flashinfer-allreduce-fusion",
|
||||
action="store_true",
|
||||
help="Enforce disable FlashInfer allreduce fusion.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--enable-aiter-allreduce-fusion",
|
||||
action="store_true",
|
||||
|
||||
Reference in New Issue
Block a user