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