Workaround of DSA performance drop on B200 + DP (#21337)
This commit is contained in:
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user