fix(attention): read per-runner kv cache dtype off model_runner (#32251)
This commit is contained in:
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user