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.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 = model_runner.kv_cache_dtype
|
||||||
|
|
||||||
self.kv_cache_dtype_str = getattr(
|
self.kv_cache_dtype_str = model_runner.kv_cache_dtype_str
|
||||||
model_runner,
|
|
||||||
"kv_cache_dtype_str",
|
|
||||||
model_runner.server_args.kv_cache_dtype,
|
|
||||||
)
|
|
||||||
self.page_size = model_runner.page_size
|
self.page_size = model_runner.page_size
|
||||||
|
|
||||||
assert self.num_heads % self.num_kv_heads == 0
|
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.token_to_kv_pool = model_runner.token_to_kv_pool
|
||||||
self.req_to_token = model_runner.req_to_token_pool.req_to_token
|
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 = model_runner.kv_cache_dtype
|
||||||
from sglang.srt.runtime_context import get_model
|
self.kv_cache_dtype_str = model_runner.kv_cache_dtype_str
|
||||||
|
|
||||||
self.kv_cache_dtype_str = get_model().kv_cache_dtype
|
|
||||||
self.kv_cache_is_mxfp8 = self.kv_cache_dtype_str == "mxfp8"
|
self.kv_cache_is_mxfp8 = self.kv_cache_dtype_str == "mxfp8"
|
||||||
self.page_size = model_runner.page_size
|
self.page_size = model_runner.page_size
|
||||||
# Static page-table width (upper bound). The device-side page-table build
|
# 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.decode_cuda_graph_metadata = {}
|
||||||
self.kv_cache_dtype = model_runner.kv_cache_dtype
|
self.kv_cache_dtype = model_runner.kv_cache_dtype
|
||||||
|
|
||||||
self.kv_cache_dtype_str = getattr(
|
self.kv_cache_dtype_str = model_runner.kv_cache_dtype_str
|
||||||
model_runner,
|
|
||||||
"kv_cache_dtype_str",
|
|
||||||
model_runner.server_args.kv_cache_dtype,
|
|
||||||
)
|
|
||||||
self.BLOCK = (
|
self.BLOCK = (
|
||||||
model_runner.model_config.block
|
model_runner.model_config.block
|
||||||
if hasattr(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.token_to_kv_pool = model_runner.token_to_kv_pool
|
||||||
self.req_to_token = model_runner.req_to_token_pool.req_to_token
|
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 = model_runner.kv_cache_dtype
|
||||||
from sglang.srt.runtime_context import get_model
|
self.kv_cache_dtype_str = model_runner.kv_cache_dtype_str
|
||||||
|
|
||||||
self.kv_cache_dtype_str = get_model().kv_cache_dtype
|
|
||||||
self.page_size = model_runner.page_size
|
self.page_size = model_runner.page_size
|
||||||
self.use_mla = model_runner.model_config.attention_arch == AttentionArch.MLA
|
self.use_mla = model_runner.model_config.attention_arch == AttentionArch.MLA
|
||||||
self.skip_prefill = skip_prefill
|
self.skip_prefill = skip_prefill
|
||||||
|
|||||||
@@ -318,6 +318,7 @@ class MockModelRunner(ModelRunner):
|
|||||||
self.device = device
|
self.device = device
|
||||||
self.dtype = dtype
|
self.dtype = dtype
|
||||||
self.kv_cache_dtype = dtype
|
self.kv_cache_dtype = dtype
|
||||||
|
self.kv_cache_dtype_str = "auto"
|
||||||
self.gpu_id = 0
|
self.gpu_id = 0
|
||||||
self.canary_manager = None
|
self.canary_manager = None
|
||||||
self.page_size = case.page_size
|
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;
|
# 656 bytes/token while the model still projects K/V in BF16;
|
||||||
# `set_mla_kv_buffer` does the quantize on the way in.
|
# `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 = torch.float8_e4m3fn if fp8_kv_cache else dtype
|
||||||
|
self.kv_cache_dtype_str = "auto"
|
||||||
# For TARGET_VERIFY / DRAFT_EXTEND, the DSA backend uses
|
# For TARGET_VERIFY / DRAFT_EXTEND, the DSA backend uses
|
||||||
# `self.speculative_num_draft_tokens` to size `seqlens_expanded`
|
# `self.speculative_num_draft_tokens` to size `seqlens_expanded`
|
||||||
# (`dsa_backend.py:482-486,510-515`). When zero, deep_gemm's
|
# (`dsa_backend.py:482-486,510-515`). When zero, deep_gemm's
|
||||||
|
|||||||
@@ -328,6 +328,7 @@ class MockDSV4ModelRunner:
|
|||||||
self.device = device
|
self.device = device
|
||||||
self.dtype = dtype
|
self.dtype = dtype
|
||||||
self.kv_cache_dtype = dtype
|
self.kv_cache_dtype = dtype
|
||||||
|
self.kv_cache_dtype_str = "auto"
|
||||||
self.gpu_id = 0
|
self.gpu_id = 0
|
||||||
self.canary_manager = None
|
self.canary_manager = None
|
||||||
self.page_size = case.page_size
|
self.page_size = case.page_size
|
||||||
|
|||||||
@@ -318,6 +318,7 @@ class DualChunkMockModelRunner(ModelRunner):
|
|||||||
self.device = device
|
self.device = device
|
||||||
self.dtype = dtype
|
self.dtype = dtype
|
||||||
self.kv_cache_dtype = dtype
|
self.kv_cache_dtype = dtype
|
||||||
|
self.kv_cache_dtype_str = "auto"
|
||||||
self.gpu_id = 0
|
self.gpu_id = 0
|
||||||
self.canary_manager = None
|
self.canary_manager = None
|
||||||
self.page_size = case.page_size
|
self.page_size = case.page_size
|
||||||
|
|||||||
@@ -212,6 +212,7 @@ class MockGDNModelRunner(ModelRunner):
|
|||||||
self.device = device
|
self.device = device
|
||||||
self.dtype = dtype
|
self.dtype = dtype
|
||||||
self.kv_cache_dtype = dtype
|
self.kv_cache_dtype = dtype
|
||||||
|
self.kv_cache_dtype_str = "auto"
|
||||||
self.gpu_id = 0
|
self.gpu_id = 0
|
||||||
self.ps = ParallelState.trivial()
|
self.ps = ParallelState.trivial()
|
||||||
self.canary_manager = None
|
self.canary_manager = None
|
||||||
|
|||||||
@@ -218,6 +218,7 @@ class MockKDAModelRunner(ModelRunner):
|
|||||||
self.device = device
|
self.device = device
|
||||||
self.dtype = dtype
|
self.dtype = dtype
|
||||||
self.kv_cache_dtype = dtype
|
self.kv_cache_dtype = dtype
|
||||||
|
self.kv_cache_dtype_str = "auto"
|
||||||
self.gpu_id = 0
|
self.gpu_id = 0
|
||||||
self.ps = ParallelState.trivial()
|
self.ps = ParallelState.trivial()
|
||||||
self.canary_manager = None
|
self.canary_manager = None
|
||||||
|
|||||||
@@ -226,6 +226,7 @@ class MockLightningModelRunner(ModelRunner):
|
|||||||
self.device = device
|
self.device = device
|
||||||
self.dtype = dtype
|
self.dtype = dtype
|
||||||
self.kv_cache_dtype = dtype
|
self.kv_cache_dtype = dtype
|
||||||
|
self.kv_cache_dtype_str = "auto"
|
||||||
self.gpu_id = 0
|
self.gpu_id = 0
|
||||||
self.ps = ParallelState.trivial()
|
self.ps = ParallelState.trivial()
|
||||||
self.canary_manager = None
|
self.canary_manager = None
|
||||||
|
|||||||
@@ -311,6 +311,7 @@ class MockMamba2ModelRunner(ModelRunner):
|
|||||||
self.device = device
|
self.device = device
|
||||||
self.dtype = dtype
|
self.dtype = dtype
|
||||||
self.kv_cache_dtype = dtype
|
self.kv_cache_dtype = dtype
|
||||||
|
self.kv_cache_dtype_str = "auto"
|
||||||
self.gpu_id = 0
|
self.gpu_id = 0
|
||||||
self.ps = ParallelState.trivial()
|
self.ps = ParallelState.trivial()
|
||||||
self.canary_manager = None
|
self.canary_manager = None
|
||||||
|
|||||||
@@ -229,6 +229,7 @@ class MockMLAModelRunner(ModelRunner):
|
|||||||
# while the model still projects K/V in bf16; `set_mla_kv_buffer`
|
# while the model still projects K/V in bf16; `set_mla_kv_buffer`
|
||||||
# does the BF16->FP8 cast on the way in.
|
# 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 = 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.gpu_id = 0
|
||||||
self.canary_manager = None
|
self.canary_manager = None
|
||||||
self.page_size = case.page_size
|
self.page_size = case.page_size
|
||||||
|
|||||||
@@ -50,13 +50,13 @@ class MockModelRunner:
|
|||||||
self.kv_cache_dtype = (
|
self.kv_cache_dtype = (
|
||||||
self.dtype
|
self.dtype
|
||||||
) # torch dtype, required by FlashAttentionBackend
|
) # torch dtype, required by FlashAttentionBackend
|
||||||
|
self.kv_cache_dtype_str = "auto"
|
||||||
|
|
||||||
# server_args is still needed for string-based config (kv_cache_dtype_str)
|
|
||||||
self.server_args = type(
|
self.server_args = type(
|
||||||
"ServerArgs",
|
"ServerArgs",
|
||||||
(),
|
(),
|
||||||
{
|
{
|
||||||
"kv_cache_dtype": "auto", # string version for kv_cache_dtype_str
|
"kv_cache_dtype": "auto",
|
||||||
"speculative_eagle_topk": None,
|
"speculative_eagle_topk": None,
|
||||||
"speculative_num_draft_tokens": 0,
|
"speculative_num_draft_tokens": 0,
|
||||||
"enable_deterministic_inference": False,
|
"enable_deterministic_inference": False,
|
||||||
|
|||||||
@@ -49,6 +49,7 @@ class MockModelRunner:
|
|||||||
},
|
},
|
||||||
)
|
)
|
||||||
self.kv_cache_dtype = self.server_args.kv_cache_dtype
|
self.kv_cache_dtype = self.server_args.kv_cache_dtype
|
||||||
|
self.kv_cache_dtype_str = "auto"
|
||||||
|
|
||||||
batch_size = 160
|
batch_size = 160
|
||||||
# Create a proper req_to_token_pool with the req_to_token attribute
|
# Create a proper req_to_token_pool with the req_to_token attribute
|
||||||
|
|||||||
Reference in New Issue
Block a user