From 2b75fed0ddebf75a80e60d99cc921611e8b52f26 Mon Sep 17 00:00:00 2001 From: Baizhou Zhang Date: Tue, 24 Mar 2026 22:21:07 -0700 Subject: [PATCH] Workaround of DSA performance drop on B200 + DP (#21337) --- python/sglang/srt/server_args.py | 16 +++++++++++----- 1 file changed, 11 insertions(+), 5 deletions(-) diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index a0f0704f3..7b855bffe 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -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: