Workaround of DSA performance drop on B200 + DP (#21337)

This commit is contained in:
Baizhou Zhang
2026-03-24 22:21:07 -07:00
committed by GitHub
parent d937d01fe6
commit 2b75fed0dd
+11 -5
View File
@@ -1386,7 +1386,7 @@ class ServerArgs:
return capture_sizes
def _set_default_nsa_kv_cache_dtype(self, major: int) -> str:
def _set_default_nsa_kv_cache_dtype(self, major: int, quantization: str) -> str:
user_set_prefill = self.nsa_prefill_backend is not None
user_set_decode = self.nsa_decode_backend is not None
@@ -1400,9 +1400,15 @@ class ServerArgs:
)
if self.kv_cache_dtype == "auto":
self.kv_cache_dtype = (
"fp8_e4m3" if (major >= 10 and self.dp_size > 1) else "bfloat16"
)
# TODO: Temporarily set default dtype on B200 as bfloat16 to avoid performance regression.
# TODO: Remove this after the performance regression is fixed. (Ref: https://github.com/sgl-project/sglang/issues/21291)
if quantization == "modelopt_fp4" and major >= 10 and self.dp_size > 1:
self.kv_cache_dtype = "fp8_e4m3"
else:
self.kv_cache_dtype = "bfloat16"
# self.kv_cache_dtype = (
# "fp8_e4m3" if (major >= 10 and self.dp_size > 1) else "bfloat16"
# )
logger.warning(
f"Setting KV cache dtype to {self.kv_cache_dtype} for DeepSeek DSA on SM{major} device."
)
@@ -1556,7 +1562,7 @@ class ServerArgs:
import torch
major, _ = torch.cuda.get_device_capability()
self._set_default_nsa_kv_cache_dtype(major)
self._set_default_nsa_kv_cache_dtype(major, self.quantization)
self._set_default_nsa_backends(self.kv_cache_dtype, major)
if self.enable_nsa_prefill_context_parallel: