fix: sync server_args.kv_cache_dtype when detecting FP8 KV cache (#18394)
This commit is contained in:
@@ -213,6 +213,13 @@ CHUNKED_PREFIX_CACHE_SUPPORTED_ATTENTION_BACKENDS = [
|
|||||||
"trtllm_mla",
|
"trtllm_mla",
|
||||||
]
|
]
|
||||||
|
|
||||||
|
TORCH_DTYPE_TO_KV_CACHE_STR = {
|
||||||
|
torch.float8_e4m3fn: "fp8_e4m3",
|
||||||
|
torch.float8_e4m3fnuz: "fp8_e4m3",
|
||||||
|
torch.float8_e5m2: "fp8_e5m2",
|
||||||
|
torch.bfloat16: "bf16",
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
def add_mla_attention_backend(backend_name):
|
def add_mla_attention_backend(backend_name):
|
||||||
if backend_name not in MLA_ATTENTION_BACKENDS:
|
if backend_name not in MLA_ATTENTION_BACKENDS:
|
||||||
@@ -1573,8 +1580,14 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
|||||||
):
|
):
|
||||||
if _is_hip:
|
if _is_hip:
|
||||||
self.kv_cache_dtype = fp8_dtype
|
self.kv_cache_dtype = fp8_dtype
|
||||||
|
self.server_args.kv_cache_dtype = TORCH_DTYPE_TO_KV_CACHE_STR[
|
||||||
|
self.kv_cache_dtype
|
||||||
|
]
|
||||||
else:
|
else:
|
||||||
self.kv_cache_dtype = torch.float8_e4m3fn
|
self.kv_cache_dtype = torch.float8_e4m3fn
|
||||||
|
self.server_args.kv_cache_dtype = TORCH_DTYPE_TO_KV_CACHE_STR[
|
||||||
|
self.kv_cache_dtype
|
||||||
|
]
|
||||||
else:
|
else:
|
||||||
self.kv_cache_dtype = self.dtype
|
self.kv_cache_dtype = self.dtype
|
||||||
elif self.server_args.kv_cache_dtype == "fp8_e5m2":
|
elif self.server_args.kv_cache_dtype == "fp8_e5m2":
|
||||||
|
|||||||
Reference in New Issue
Block a user