fix(attention): read per-runner kv cache dtype off model_runner (#32251)

This commit is contained in:
Cheng Wan
2026-07-23 20:08:57 -07:00
committed by GitHub
parent bd3f6a7935
commit eac7c7d7cd
15 changed files with 16 additions and 18 deletions
@@ -126,11 +126,7 @@ class DualChunkFlashAttentionBackend(AttentionBackend):
self.req_to_token = model_runner.req_to_token_pool.req_to_token
self.kv_cache_dtype = model_runner.kv_cache_dtype
self.kv_cache_dtype_str = getattr(
model_runner,
"kv_cache_dtype_str",
model_runner.server_args.kv_cache_dtype,
)
self.kv_cache_dtype_str = model_runner.kv_cache_dtype_str
self.page_size = model_runner.page_size
assert self.num_heads % self.num_kv_heads == 0
@@ -166,9 +166,7 @@ class FlashAttentionBackend(AttentionBackend):
self.token_to_kv_pool = model_runner.token_to_kv_pool
self.req_to_token = model_runner.req_to_token_pool.req_to_token
self.kv_cache_dtype = model_runner.kv_cache_dtype
from sglang.srt.runtime_context import get_model
self.kv_cache_dtype_str = get_model().kv_cache_dtype
self.kv_cache_dtype_str = model_runner.kv_cache_dtype_str
self.kv_cache_is_mxfp8 = self.kv_cache_dtype_str == "mxfp8"
self.page_size = model_runner.page_size
# Static page-table width (upper bound). The device-side page-table build
@@ -65,11 +65,7 @@ class LightningAttentionBackend(MambaAttnBackendBase):
self.decode_cuda_graph_metadata = {}
self.kv_cache_dtype = model_runner.kv_cache_dtype
self.kv_cache_dtype_str = getattr(
model_runner,
"kv_cache_dtype_str",
model_runner.server_args.kv_cache_dtype,
)
self.kv_cache_dtype_str = model_runner.kv_cache_dtype_str
self.BLOCK = (
model_runner.model_config.block
if hasattr(model_runner.model_config, "block")
@@ -69,9 +69,7 @@ class XPUAttentionBackend(AttentionBackend):
self.token_to_kv_pool = model_runner.token_to_kv_pool
self.req_to_token = model_runner.req_to_token_pool.req_to_token
self.kv_cache_dtype = model_runner.kv_cache_dtype
from sglang.srt.runtime_context import get_model
self.kv_cache_dtype_str = get_model().kv_cache_dtype
self.kv_cache_dtype_str = model_runner.kv_cache_dtype_str
self.page_size = model_runner.page_size
self.use_mla = model_runner.model_config.attention_arch == AttentionArch.MLA
self.skip_prefill = skip_prefill
@@ -318,6 +318,7 @@ class MockModelRunner(ModelRunner):
self.device = device
self.dtype = dtype
self.kv_cache_dtype = dtype
self.kv_cache_dtype_str = "auto"
self.gpu_id = 0
self.canary_manager = None
self.page_size = case.page_size
@@ -289,6 +289,7 @@ class DSAMockModelRunner(ModelRunner):
# 656 bytes/token while the model still projects K/V in BF16;
# `set_mla_kv_buffer` does the quantize on the way in.
self.kv_cache_dtype = torch.float8_e4m3fn if fp8_kv_cache else dtype
self.kv_cache_dtype_str = "auto"
# For TARGET_VERIFY / DRAFT_EXTEND, the DSA backend uses
# `self.speculative_num_draft_tokens` to size `seqlens_expanded`
# (`dsa_backend.py:482-486,510-515`). When zero, deep_gemm's
@@ -328,6 +328,7 @@ class MockDSV4ModelRunner:
self.device = device
self.dtype = dtype
self.kv_cache_dtype = dtype
self.kv_cache_dtype_str = "auto"
self.gpu_id = 0
self.canary_manager = None
self.page_size = case.page_size
@@ -318,6 +318,7 @@ class DualChunkMockModelRunner(ModelRunner):
self.device = device
self.dtype = dtype
self.kv_cache_dtype = dtype
self.kv_cache_dtype_str = "auto"
self.gpu_id = 0
self.canary_manager = None
self.page_size = case.page_size
@@ -212,6 +212,7 @@ class MockGDNModelRunner(ModelRunner):
self.device = device
self.dtype = dtype
self.kv_cache_dtype = dtype
self.kv_cache_dtype_str = "auto"
self.gpu_id = 0
self.ps = ParallelState.trivial()
self.canary_manager = None
@@ -218,6 +218,7 @@ class MockKDAModelRunner(ModelRunner):
self.device = device
self.dtype = dtype
self.kv_cache_dtype = dtype
self.kv_cache_dtype_str = "auto"
self.gpu_id = 0
self.ps = ParallelState.trivial()
self.canary_manager = None
@@ -226,6 +226,7 @@ class MockLightningModelRunner(ModelRunner):
self.device = device
self.dtype = dtype
self.kv_cache_dtype = dtype
self.kv_cache_dtype_str = "auto"
self.gpu_id = 0
self.ps = ParallelState.trivial()
self.canary_manager = None
@@ -311,6 +311,7 @@ class MockMamba2ModelRunner(ModelRunner):
self.device = device
self.dtype = dtype
self.kv_cache_dtype = dtype
self.kv_cache_dtype_str = "auto"
self.gpu_id = 0
self.ps = ParallelState.trivial()
self.canary_manager = None
@@ -229,6 +229,7 @@ class MockMLAModelRunner(ModelRunner):
# while the model still projects K/V in bf16; `set_mla_kv_buffer`
# does the BF16->FP8 cast on the way in.
self.kv_cache_dtype = torch.float8_e4m3fn if fp8_kv_cache else dtype
self.kv_cache_dtype_str = "fp8_e4m3" if fp8_kv_cache else "auto"
self.gpu_id = 0
self.canary_manager = None
self.page_size = case.page_size