[Refactor] Generalize DeepSeek V4 compressed pool management (#38954)
This commit is contained in:
@@ -321,8 +321,8 @@ class DSV4NPUTokenToKVPool(DeepSeekV4TokenToKVPool):
|
|||||||
return self.get_indexer_compress_states(layer_id)
|
return self.get_indexer_compress_states(layer_id)
|
||||||
return self.get_attention_compress_states(layer_id)
|
return self.get_attention_compress_states(layer_id)
|
||||||
|
|
||||||
def _make_attn_state_pool(
|
def _make_compress_state_pool(
|
||||||
self, ratio: int, enable_memory_saver: bool
|
self, ratio: int, *, head_dim: int, enable_memory_saver: bool
|
||||||
) -> NPUCompressStatePool:
|
) -> NPUCompressStatePool:
|
||||||
# ONLINE_C128 (CUDA-only) collapses the c128 ring to size 1; the NPU fused
|
# ONLINE_C128 (CUDA-only) collapses the c128 ring to size 1; the NPU fused
|
||||||
# compressor has no online mode, so assert the config mismatch early.
|
# compressor has no online mode, so assert the config mismatch early.
|
||||||
@@ -330,46 +330,26 @@ class DSV4NPUTokenToKVPool(DeepSeekV4TokenToKVPool):
|
|||||||
"SGLANG_OPT_USE_ONLINE_COMPRESS is incompatible with the "
|
"SGLANG_OPT_USE_ONLINE_COMPRESS is incompatible with the "
|
||||||
"NPU fused compressor (no online mode in the kernel)."
|
"NPU fused compressor (no online mode in the kernel)."
|
||||||
)
|
)
|
||||||
|
config = self.compressed_pool_configs[ratio]
|
||||||
ring_size = self.get_ring_size(ratio)
|
ring_size = self.get_ring_size(ratio)
|
||||||
# A5 cache_mode=2 addresses one ring bank per request. The A3
|
# A5 cache_mode=2 addresses one ring bank per request. The A3
|
||||||
# explicit-location path can share the smaller flat pool, but the A5
|
# explicit-location path can share the smaller flat pool, but the A5
|
||||||
# cycle ABI needs enough physical banks for every req_pool_idx.
|
# cycle ABI needs enough physical banks for every req_pool_idx.
|
||||||
size = self._state_pool_size(ratio)
|
size = config.state_size
|
||||||
if is_npu_arch35():
|
if is_npu_arch35():
|
||||||
size = max(size, self.num_req_slots * ring_size)
|
size = max(size, self.num_req_slots * ring_size)
|
||||||
return NPUCompressStatePool(
|
return NPUCompressStatePool(
|
||||||
size=size,
|
size=size,
|
||||||
ring_size=ring_size,
|
ring_size=ring_size,
|
||||||
overlap=ratio == 4,
|
overlap=ratio == 4,
|
||||||
head_dim=self.qk_nope_head_dim + self.qk_rope_head_dim,
|
head_dim=head_dim,
|
||||||
dtype=self.c4_state_dtype if ratio == 4 else self.c128_state_dtype,
|
dtype=config.state_dtype,
|
||||||
device=self.device,
|
device=self.device,
|
||||||
enable_memory_saver=enable_memory_saver,
|
enable_memory_saver=enable_memory_saver,
|
||||||
ratio=ratio,
|
ratio=ratio,
|
||||||
swa_page_size=self.swa_page_size,
|
swa_page_size=self.swa_page_size,
|
||||||
)
|
)
|
||||||
|
|
||||||
def _make_indexer_state_pool(
|
|
||||||
self, ratio: int, enable_memory_saver: bool
|
|
||||||
) -> NPUCompressStatePool:
|
|
||||||
# c4 indexer shares the c4 state pool size budget but has its own
|
|
||||||
# slot_dim (indexer_head_dim vs attention head_dim).
|
|
||||||
ring_size = self.get_ring_size(ratio)
|
|
||||||
size = self.c4_state_pool_size
|
|
||||||
if is_npu_arch35():
|
|
||||||
size = max(size, self.num_req_slots * ring_size)
|
|
||||||
return NPUCompressStatePool(
|
|
||||||
size=size,
|
|
||||||
ring_size=ring_size,
|
|
||||||
overlap=ratio == 4,
|
|
||||||
head_dim=self.indexer_head_dim,
|
|
||||||
device=self.device,
|
|
||||||
dtype=self.c4_state_dtype,
|
|
||||||
enable_memory_saver=enable_memory_saver,
|
|
||||||
ratio=ratio,
|
|
||||||
swa_page_size=self.swa_page_size,
|
|
||||||
)
|
|
||||||
|
|
||||||
def _make_indexer_pool(
|
def _make_indexer_pool(
|
||||||
self,
|
self,
|
||||||
size: int,
|
size: int,
|
||||||
@@ -394,10 +374,11 @@ class DSV4NPUTokenToKVPool(DeepSeekV4TokenToKVPool):
|
|||||||
|
|
||||||
def get_contiguous_buf_infos(self) -> Tuple[List[int], List[int], List[int]]:
|
def get_contiguous_buf_infos(self) -> Tuple[List[int], List[int], List[int]]:
|
||||||
"""Main PD buffers addressed by the full KV page id."""
|
"""Main PD buffers addressed by the full KV page id."""
|
||||||
|
indexer_pool = self._indexer_pool(4)
|
||||||
buffers = (
|
buffers = (
|
||||||
self.c4_kv_pool.kv_buffer
|
self.c4_kv_pool.kv_buffer
|
||||||
+ self.c4_indexer_kv_pool.index_k_buffer
|
+ indexer_pool.index_k_buffer
|
||||||
+ self.c4_indexer_kv_pool.index_scale_buffer
|
+ indexer_pool.index_scale_buffer
|
||||||
)
|
)
|
||||||
return (
|
return (
|
||||||
[buf.data_ptr() for buf in buffers],
|
[buf.data_ptr() for buf in buffers],
|
||||||
@@ -530,16 +511,15 @@ class DSV4NPUTokenToKVPool(DeepSeekV4TokenToKVPool):
|
|||||||
``torch.ops.custom.npu_quant_lightning_indexer`` consumes.
|
``torch.ops.custom.npu_quant_lightning_indexer`` consumes.
|
||||||
"""
|
"""
|
||||||
item = self.layer_mapping[layer_id]
|
item = self.layer_mapping[layer_id]
|
||||||
if item.compress_ratio == 4:
|
if item.compress_ratio == 0:
|
||||||
if from_indexer:
|
|
||||||
kv = self.c4_indexer_kv_pool.get_index_k(item.compress_layer_id)
|
|
||||||
else:
|
|
||||||
kv = self.c4_kv_pool.kv_buffer[item.compress_layer_id]
|
|
||||||
elif item.compress_ratio == 128:
|
|
||||||
assert not from_indexer, "c128 has no indexer pool"
|
|
||||||
kv = self.c128_kv_pool.kv_buffer[item.compress_layer_id]
|
|
||||||
else:
|
|
||||||
return None
|
return None
|
||||||
|
if from_indexer:
|
||||||
|
indexer_pool = self._indexer_pool(item.compress_ratio)
|
||||||
|
kv = indexer_pool.get_index_k(item.compress_layer_id)
|
||||||
|
else:
|
||||||
|
compress_pool = item.compress_kv_pool
|
||||||
|
assert compress_pool is not None, "Missing compressed KV pool"
|
||||||
|
kv = compress_pool.kv_buffer[item.compress_layer_id]
|
||||||
if loc is not None:
|
if loc is not None:
|
||||||
kv = kv.flatten(0, 1)[loc]
|
kv = kv.flatten(0, 1)[loc]
|
||||||
return kv
|
return kv
|
||||||
@@ -651,21 +631,17 @@ class DSV4NPUTokenToKVPool(DeepSeekV4TokenToKVPool):
|
|||||||
ratio, compress_layer_id, _ = self.layer_mapping[layer_id]
|
ratio, compress_layer_id, _ = self.layer_mapping[layer_id]
|
||||||
device_type = kv.device.type
|
device_type = kv.device.type
|
||||||
if from_indexer:
|
if from_indexer:
|
||||||
assert ratio == 4, f"indexer only on c4 layers, got ratio={ratio}"
|
indexer_pool = self._indexer_pool(ratio)
|
||||||
if device_type == "npu":
|
if device_type == "npu":
|
||||||
assert self.c4_indexer_kv_pool.has_npu_storage, (
|
assert indexer_pool.has_npu_storage, (
|
||||||
"NPU index buffers not allocated — pool was init'd on CUDA?"
|
"NPU index buffers not allocated — pool was init'd on CUDA?"
|
||||||
)
|
)
|
||||||
self.c4_indexer_kv_pool.set_index_k_scale(
|
indexer_pool.set_index_k_scale(compress_layer_id, loc, kv, kv_scale)
|
||||||
compress_layer_id, loc, kv, kv_scale
|
|
||||||
)
|
|
||||||
return
|
return
|
||||||
if kv_scale is None:
|
if kv_scale is None:
|
||||||
self.c4_indexer_kv_pool.set_index_fused(compress_layer_id, loc, kv)
|
indexer_pool.set_index_fused(compress_layer_id, loc, kv)
|
||||||
return
|
return
|
||||||
self.c4_indexer_kv_pool.set_index_k_scale_buffer(
|
indexer_pool.set_index_k_scale_buffer(compress_layer_id, loc, kv, kv_scale)
|
||||||
compress_layer_id, loc, kv, kv_scale
|
|
||||||
)
|
|
||||||
return
|
return
|
||||||
compress_pool = self.c4_kv_pool if ratio == 4 else self.c128_kv_pool
|
compress_pool = self.c4_kv_pool if ratio == 4 else self.c128_kv_pool
|
||||||
if device_type == "npu":
|
if device_type == "npu":
|
||||||
@@ -690,5 +666,6 @@ class DSV4NPUTokenToKVPool(DeepSeekV4TokenToKVPool):
|
|||||||
) -> torch.Tensor:
|
) -> torch.Tensor:
|
||||||
# The indexer scale is fp16 on pre-A5 parts and fp32 on A5.
|
# The indexer scale is fp16 on pre-A5 parts and fp32 on A5.
|
||||||
assert from_indexer, "only indexer compress pool has dequant scale"
|
assert from_indexer, "only indexer compress pool has dequant scale"
|
||||||
compress_layer_id = self.layer_mapping[layer_id].compress_layer_id
|
item = self.layer_mapping[layer_id]
|
||||||
return self.c4_indexer_kv_pool.get_index_scale(compress_layer_id)
|
indexer_pool = self._indexer_pool(item.compress_ratio)
|
||||||
|
return indexer_pool.get_index_scale(item.compress_layer_id)
|
||||||
|
|||||||
@@ -778,6 +778,7 @@ class DeepseekV4AttnBackend(
|
|||||||
):
|
):
|
||||||
return PagedIndexerMetadata(
|
return PagedIndexerMetadata(
|
||||||
page_size=self.page_size,
|
page_size=self.page_size,
|
||||||
|
compressed_page_size=self.token_to_kv_pool.get_index_k_page_size(),
|
||||||
page_table=core_attn_metadata.page_table,
|
page_table=core_attn_metadata.page_table,
|
||||||
compressed_seq_lens=core_attn_metadata.c4_topk_lengths_raw,
|
compressed_seq_lens=core_attn_metadata.c4_topk_lengths_raw,
|
||||||
use_topk_v2=self.dsa_topk_backend.should_use_topk_v2() and not _is_xpu,
|
use_topk_v2=self.dsa_topk_backend.should_use_topk_v2() and not _is_xpu,
|
||||||
@@ -1482,8 +1483,6 @@ class DeepseekV4AttnBackend(
|
|||||||
seq_lens_cpu_list, extend_seq_lens_cpu, strict=True
|
seq_lens_cpu_list, extend_seq_lens_cpu, strict=True
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
# ``swa_window_size`` on the pool is its storage page size, not the
|
|
||||||
# model's SWA window, so pass both explicitly.
|
|
||||||
return SparsePrefillChunkCache.build(
|
return SparsePrefillChunkCache.build(
|
||||||
seq_lens=forward_batch.seq_lens.to(torch.int32),
|
seq_lens=forward_batch.seq_lens.to(torch.int32),
|
||||||
extend_seq_lens=forward_batch.extend_seq_lens.to(torch.int32),
|
extend_seq_lens=forward_batch.extend_seq_lens.to(torch.int32),
|
||||||
@@ -1491,7 +1490,7 @@ class DeepseekV4AttnBackend(
|
|||||||
req_to_token=self.req_to_token,
|
req_to_token=self.req_to_token,
|
||||||
full_to_swa=self.token_to_kv_pool.full_to_swa_index_mapping,
|
full_to_swa=self.token_to_kv_pool.full_to_swa_index_mapping,
|
||||||
swa_window_size=SWA_WINDOW,
|
swa_window_size=SWA_WINDOW,
|
||||||
swa_page_size=self.token_to_kv_pool.swa_window_size,
|
swa_page_size=self.token_to_kv_pool.swa_page_size,
|
||||||
num_qo_tokens=num_qo_tokens,
|
num_qo_tokens=num_qo_tokens,
|
||||||
max_seq_len=max(seq_lens_cpu_list),
|
max_seq_len=max(seq_lens_cpu_list),
|
||||||
total_swa=total_swa,
|
total_swa=total_swa,
|
||||||
@@ -1772,23 +1771,20 @@ class DeepseekV4AttnBackend(
|
|||||||
extra_indices = core_attn_metadata.c128_page_indices
|
extra_indices = core_attn_metadata.c128_page_indices
|
||||||
extra_topk_lengths = core_attn_metadata.c128_topk_lengths_clamp1
|
extra_topk_lengths = core_attn_metadata.c128_topk_lengths_clamp1
|
||||||
|
|
||||||
swa_window_size = token_to_kv_pool.swa_window_size
|
swa_page_size = token_to_kv_pool.swa_page_size
|
||||||
assert swa_k_cache.ndim == 2
|
assert swa_k_cache.ndim == 2
|
||||||
k_cache_total_dim = token_to_kv_pool.swa_kv_pool.kv_cache_total_dim
|
k_cache_total_dim = token_to_kv_pool.swa_kv_pool.kv_cache_total_dim
|
||||||
swa_k_cache = swa_k_cache[:, : swa_window_size * k_cache_total_dim].view(
|
swa_k_cache = swa_k_cache[:, : swa_page_size * k_cache_total_dim].view(
|
||||||
swa_k_cache.shape[0], swa_window_size, 1, k_cache_total_dim
|
swa_k_cache.shape[0], swa_page_size, 1, k_cache_total_dim
|
||||||
)
|
)
|
||||||
|
|
||||||
if extra_k_cache is not None:
|
if extra_k_cache is not None:
|
||||||
page_sizes = {
|
extra_page_size = token_to_kv_pool.get_extra_key_page_size(layer_id)
|
||||||
4: token_to_kv_pool.page_size // 4,
|
|
||||||
128: token_to_kv_pool.page_size // 128,
|
|
||||||
}
|
|
||||||
extra_k_cache = extra_k_cache[
|
extra_k_cache = extra_k_cache[
|
||||||
:, : page_sizes[compress_ratio] * k_cache_total_dim
|
:, : extra_page_size * k_cache_total_dim
|
||||||
].view(
|
].view(
|
||||||
extra_k_cache.shape[0],
|
extra_k_cache.shape[0],
|
||||||
page_sizes[compress_ratio],
|
extra_page_size,
|
||||||
1,
|
1,
|
||||||
k_cache_total_dim,
|
k_cache_total_dim,
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -507,6 +507,7 @@ class DeepseekV4HipRadixBackend(
|
|||||||
def init_forward_metadata_indexer(self, core_attn_metadata: DSV4AttnMetadata):
|
def init_forward_metadata_indexer(self, core_attn_metadata: DSV4AttnMetadata):
|
||||||
return PagedIndexerMetadata(
|
return PagedIndexerMetadata(
|
||||||
page_size=self.page_size,
|
page_size=self.page_size,
|
||||||
|
compressed_page_size=self.token_to_kv_pool.get_index_k_page_size(),
|
||||||
page_table=core_attn_metadata.page_table,
|
page_table=core_attn_metadata.page_table,
|
||||||
compressed_seq_lens=core_attn_metadata.c4_topk_lengths_raw,
|
compressed_seq_lens=core_attn_metadata.c4_topk_lengths_raw,
|
||||||
use_topk_v2=self.dsa_topk_backend.should_use_topk_v2(),
|
use_topk_v2=self.dsa_topk_backend.should_use_topk_v2(),
|
||||||
@@ -1601,23 +1602,20 @@ class DeepseekV4HipRadixBackend(
|
|||||||
extra_indices = core_attn_metadata.c128_page_indices
|
extra_indices = core_attn_metadata.c128_page_indices
|
||||||
extra_topk_lengths = core_attn_metadata.c128_topk_lengths_clamp1
|
extra_topk_lengths = core_attn_metadata.c128_topk_lengths_clamp1
|
||||||
|
|
||||||
swa_window_size = token_to_kv_pool.swa_window_size
|
swa_page_size = token_to_kv_pool.swa_page_size
|
||||||
assert swa_k_cache.ndim == 2
|
assert swa_k_cache.ndim == 2
|
||||||
k_cache_total_dim = token_to_kv_pool.swa_kv_pool.kv_cache_total_dim
|
k_cache_total_dim = token_to_kv_pool.swa_kv_pool.kv_cache_total_dim
|
||||||
swa_k_cache = swa_k_cache[:, : swa_window_size * k_cache_total_dim].view(
|
swa_k_cache = swa_k_cache[:, : swa_page_size * k_cache_total_dim].view(
|
||||||
swa_k_cache.shape[0], swa_window_size, 1, k_cache_total_dim
|
swa_k_cache.shape[0], swa_page_size, 1, k_cache_total_dim
|
||||||
)
|
)
|
||||||
|
|
||||||
if extra_k_cache is not None:
|
if extra_k_cache is not None:
|
||||||
page_sizes = {
|
extra_page_size = token_to_kv_pool.get_extra_key_page_size(layer_id)
|
||||||
4: token_to_kv_pool.page_size // 4,
|
|
||||||
128: token_to_kv_pool.page_size // 128,
|
|
||||||
}
|
|
||||||
extra_k_cache = extra_k_cache[
|
extra_k_cache = extra_k_cache[
|
||||||
:, : page_sizes[compress_ratio] * k_cache_total_dim
|
:, : extra_page_size * k_cache_total_dim
|
||||||
].view(
|
].view(
|
||||||
extra_k_cache.shape[0],
|
extra_k_cache.shape[0],
|
||||||
page_sizes[compress_ratio],
|
extra_page_size,
|
||||||
1,
|
1,
|
||||||
k_cache_total_dim,
|
k_cache_total_dim,
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -265,7 +265,7 @@ class CompressorBackendMixin:
|
|||||||
bf16_store = False
|
bf16_store = False
|
||||||
kv_scale_cache = None
|
kv_scale_cache = None
|
||||||
if compressor.is_in_indexer:
|
if compressor.is_in_indexer:
|
||||||
page_size = token_to_kv_pool.get_index_k_page_size()
|
page_size = token_to_kv_pool.get_index_k_page_size(compressor.ratio)
|
||||||
if use_hip_fp4:
|
if use_hip_fp4:
|
||||||
kv_cache = token_to_kv_pool.get_index_k_fp4_payload_buffer(layer_id)
|
kv_cache = token_to_kv_pool.get_index_k_fp4_payload_buffer(layer_id)
|
||||||
kv_scale_cache = token_to_kv_pool.get_index_k_fp4_scale_buffer(layer_id)
|
kv_scale_cache = token_to_kv_pool.get_index_k_fp4_scale_buffer(layer_id)
|
||||||
|
|||||||
@@ -2,7 +2,7 @@ from __future__ import annotations
|
|||||||
|
|
||||||
import warnings
|
import warnings
|
||||||
from dataclasses import dataclass, field, fields
|
from dataclasses import dataclass, field, fields
|
||||||
from typing import TYPE_CHECKING, Any, List, Optional
|
from typing import Any, List, Optional
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
@@ -11,10 +11,6 @@ from sglang.srt.utils import is_hip, is_sm120_supported, is_xpu
|
|||||||
|
|
||||||
_IS_SM120 = is_sm120_supported()
|
_IS_SM120 = is_sm120_supported()
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
|
||||||
pass
|
|
||||||
|
|
||||||
|
|
||||||
"""
|
"""
|
||||||
Some comments on the common terms used in DeepSeekV4Backend:
|
Some comments on the common terms used in DeepSeekV4Backend:
|
||||||
|
|
||||||
@@ -45,7 +41,7 @@ positions:
|
|||||||
|
|
||||||
Some other notes:
|
Some other notes:
|
||||||
c4_ / c128_: means "compressed by 4" / "compressed by 128".
|
c4_ / c128_: means "compressed by 4" / "compressed by 128".
|
||||||
compressed_page_size: page_size // 4
|
compressed_page_size: physical indexer pool page size
|
||||||
compressed_seq_lens: seq_lens // 4, but bounded by at least 1, due to flash_mla requirement.
|
compressed_seq_lens: seq_lens // 4, but bounded by at least 1, due to flash_mla requirement.
|
||||||
c4_sparse: means "compressed by 4" but only attend to top-512 tokens.
|
c4_sparse: means "compressed by 4" but only attend to top-512 tokens.
|
||||||
all related length will be clipped to 512.
|
all related length will be clipped to 512.
|
||||||
@@ -114,6 +110,7 @@ class NonPagedIndexerPlan:
|
|||||||
@dataclass
|
@dataclass
|
||||||
class PagedIndexerMetadata:
|
class PagedIndexerMetadata:
|
||||||
page_size: int
|
page_size: int
|
||||||
|
compressed_page_size: int
|
||||||
page_table: torch.Tensor
|
page_table: torch.Tensor
|
||||||
compressed_seq_lens: torch.Tensor
|
compressed_seq_lens: torch.Tensor
|
||||||
use_topk_v2: bool
|
use_topk_v2: bool
|
||||||
@@ -177,10 +174,6 @@ class PagedIndexerMetadata:
|
|||||||
|
|
||||||
assert self.page_size == 256, "the system hardcodes page_size=256"
|
assert self.page_size == 256, "the system hardcodes page_size=256"
|
||||||
|
|
||||||
@property
|
|
||||||
def compressed_page_size(self) -> int:
|
|
||||||
return self.page_size // 4
|
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def max_seq_len(self) -> int:
|
def max_seq_len(self) -> int:
|
||||||
return self.page_table.shape[1] * self.page_size
|
return self.page_table.shape[1] * self.page_size
|
||||||
@@ -202,6 +195,7 @@ class PagedIndexerMetadata:
|
|||||||
dst=self,
|
dst=self,
|
||||||
check_eq_fields=[
|
check_eq_fields=[
|
||||||
"page_size",
|
"page_size",
|
||||||
|
"compressed_page_size",
|
||||||
"force_deep_gemm_metadata",
|
"force_deep_gemm_metadata",
|
||||||
"use_prefill_cuda_graph",
|
"use_prefill_cuda_graph",
|
||||||
"use_topk_v2",
|
"use_topk_v2",
|
||||||
|
|||||||
@@ -2,7 +2,7 @@ from __future__ import annotations
|
|||||||
|
|
||||||
import logging
|
import logging
|
||||||
from contextlib import nullcontext
|
from contextlib import nullcontext
|
||||||
from typing import List, Literal, NamedTuple, Optional, Sequence, Tuple
|
from typing import List, NamedTuple, Optional, Sequence, Tuple
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
@@ -498,8 +498,15 @@ class DeepSeekV4IndexerPool(KVCache):
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class _CompressedPoolConfig(NamedTuple):
|
||||||
|
kv_size: int
|
||||||
|
state_size: int
|
||||||
|
state_dtype: torch.dtype
|
||||||
|
indexer_size: Optional[int] = None
|
||||||
|
|
||||||
|
|
||||||
class DeepSeekV4LayerItem(NamedTuple):
|
class DeepSeekV4LayerItem(NamedTuple):
|
||||||
compress_ratio: Literal[0, 4, 128]
|
compress_ratio: int
|
||||||
compress_layer_id: int
|
compress_layer_id: int
|
||||||
compress_kv_pool: Optional[DeepSeekV4SingleKVPool] = None
|
compress_kv_pool: Optional[DeepSeekV4SingleKVPool] = None
|
||||||
|
|
||||||
@@ -623,7 +630,6 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool):
|
|||||||
self.num_req_slots = (
|
self.num_req_slots = (
|
||||||
num_req_slots if num_req_slots is not None else max_num_reqs + 1
|
num_req_slots if num_req_slots is not None else max_num_reqs + 1
|
||||||
)
|
)
|
||||||
self.c4_size = c4_size
|
|
||||||
self.c4_logical_size = c4_logical_size
|
self.c4_logical_size = c4_logical_size
|
||||||
self.c128_size = c128_size
|
self.c128_size = c128_size
|
||||||
from sglang.kernels.ops.attention.dsv4.unified_kv_kernels.env_gate import (
|
from sglang.kernels.ops.attention.dsv4.unified_kv_kernels.env_gate import (
|
||||||
@@ -642,7 +648,6 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool):
|
|||||||
# so the caller-supplied, SWA-scaled size does not apply here.
|
# so the caller-supplied, SWA-scaled size does not apply here.
|
||||||
c4_state_pool_size = self.num_req_slots * c4_ring_size
|
c4_state_pool_size = self.num_req_slots * c4_ring_size
|
||||||
# Non-unified (fp8) keeps the caller-supplied, SWA-addressed size.
|
# Non-unified (fp8) keeps the caller-supplied, SWA-addressed size.
|
||||||
self.c4_state_pool_size = c4_state_pool_size
|
|
||||||
c128_ring_size = self.get_ring_size(128)
|
c128_ring_size = self.get_ring_size(128)
|
||||||
if ONLINE_C128:
|
if ONLINE_C128:
|
||||||
# Request-scoped C128 state must also cover PD preallocation slots.
|
# Request-scoped C128 state must also cover PD preallocation slots.
|
||||||
@@ -652,9 +657,19 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool):
|
|||||||
c128_state_pool_size = max(
|
c128_state_pool_size = max(
|
||||||
c128_state_pool_size, self.num_req_slots * c128_ring_size
|
c128_state_pool_size, self.num_req_slots * c128_ring_size
|
||||||
)
|
)
|
||||||
self.c128_state_pool_size = c128_state_pool_size
|
self.compressed_pool_configs = {
|
||||||
self.c4_state_dtype = c4_state_dtype
|
4: _CompressedPoolConfig(
|
||||||
self.c128_state_dtype = c128_state_dtype
|
kv_size=c4_size,
|
||||||
|
state_size=c4_state_pool_size,
|
||||||
|
state_dtype=c4_state_dtype,
|
||||||
|
indexer_size=c4_logical_size,
|
||||||
|
),
|
||||||
|
128: _CompressedPoolConfig(
|
||||||
|
kv_size=c128_size,
|
||||||
|
state_size=c128_state_pool_size,
|
||||||
|
state_dtype=c128_state_dtype,
|
||||||
|
),
|
||||||
|
}
|
||||||
self.compression_ratios = compression_ratios
|
self.compression_ratios = compression_ratios
|
||||||
self.online_mtp_max_draft_tokens = online_mtp_max_draft_tokens
|
self.online_mtp_max_draft_tokens = online_mtp_max_draft_tokens
|
||||||
self.online_c128_state_num_req_slots = c128_state_pool_size
|
self.online_c128_state_num_req_slots = c128_state_pool_size
|
||||||
@@ -681,24 +696,17 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool):
|
|||||||
self.sliding_window = sliding_window
|
self.sliding_window = sliding_window
|
||||||
|
|
||||||
self.swa_size = swa_size
|
self.swa_size = swa_size
|
||||||
self.swa_window_size = swa_page_size
|
|
||||||
self.swa_page_size = swa_page_size
|
self.swa_page_size = swa_page_size
|
||||||
self.scale_pad = 1
|
|
||||||
|
|
||||||
self.qk_nope_head_dim = qk_nope_head_dim
|
self.qk_nope_head_dim = qk_nope_head_dim
|
||||||
self.qk_rope_head_dim = qk_rope_head_dim
|
self.qk_rope_head_dim = qk_rope_head_dim
|
||||||
self.indexer_head_dim = indexer_head_dim
|
self.indexer_head_dim = indexer_head_dim
|
||||||
|
|
||||||
stage_layer_num = len(stage_ratios)
|
stage_layer_num = len(stage_ratios)
|
||||||
c4_layer_num = sum(1 for r in stage_ratios if r == 4)
|
kv_pool_cls: type = DeepSeekV4SingleKVPool
|
||||||
c128_layer_num = sum(1 for r in stage_ratios if r == 128)
|
|
||||||
c4_page_size = page_size // 4
|
|
||||||
c128_page_size = page_size // 128
|
|
||||||
|
|
||||||
if self._unified_kv:
|
if self._unified_kv:
|
||||||
self.swa_kv_pool = None
|
self.swa_kv_pool = None
|
||||||
self.c4_kv_pool = None
|
|
||||||
self.c128_kv_pool = None
|
|
||||||
swa_ring_size = get_swa_ring_size(
|
swa_ring_size = get_swa_ring_size(
|
||||||
self.sliding_window, get_spec().speculative_algorithm is not None
|
self.sliding_window, get_spec().speculative_algorithm is not None
|
||||||
)
|
)
|
||||||
@@ -721,7 +729,6 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool):
|
|||||||
self.swa_req_ring_size = self.unified_swa_ring_size
|
self.swa_req_ring_size = self.unified_swa_ring_size
|
||||||
else:
|
else:
|
||||||
self.unified_kv_pool = None
|
self.unified_kv_pool = None
|
||||||
kv_pool_cls: type = DeepSeekV4SingleKVPool
|
|
||||||
if self.uniform_fp8:
|
if self.uniform_fp8:
|
||||||
assert dtype == torch.float8_e4m3fn, (
|
assert dtype == torch.float8_e4m3fn, (
|
||||||
"--dsv4-attn-backend trtllm requires "
|
"--dsv4-attn-backend trtllm requires "
|
||||||
@@ -739,43 +746,14 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool):
|
|||||||
cls=kv_pool_cls,
|
cls=kv_pool_cls,
|
||||||
)
|
)
|
||||||
|
|
||||||
c4_kv_pool_type = kv_pool_cls
|
self._init_compressed_pools(
|
||||||
if enable_hisparse:
|
stage_ratios=stage_ratios,
|
||||||
assert not self.uniform_fp8, (
|
page_size=page_size,
|
||||||
"enable_hisparse is not supported with --dsv4-attn-backend trtllm."
|
dtype=dtype,
|
||||||
)
|
device=device,
|
||||||
c4_kv_pool_type = HiSparseC4DevicePool
|
enable_memory_saver=enable_memory_saver,
|
||||||
self.c4_kv_pool = self._make_kv_pool(
|
enable_hisparse=enable_hisparse,
|
||||||
size=c4_size,
|
kv_pool_cls=kv_pool_cls,
|
||||||
page_size=c4_page_size,
|
|
||||||
dtype=dtype,
|
|
||||||
layer_num=c4_layer_num,
|
|
||||||
device=device,
|
|
||||||
enable_memory_saver=enable_memory_saver,
|
|
||||||
global_page_size=page_size,
|
|
||||||
cls=c4_kv_pool_type,
|
|
||||||
)
|
|
||||||
|
|
||||||
self.c128_kv_pool = self._make_kv_pool(
|
|
||||||
size=c128_size,
|
|
||||||
page_size=c128_page_size,
|
|
||||||
dtype=dtype,
|
|
||||||
layer_num=c128_layer_num,
|
|
||||||
device=device,
|
|
||||||
enable_memory_saver=enable_memory_saver,
|
|
||||||
global_page_size=page_size,
|
|
||||||
cls=kv_pool_cls,
|
|
||||||
)
|
|
||||||
|
|
||||||
indexer_size = self.c4_logical_size
|
|
||||||
self.c4_indexer_kv_pool = self._make_indexer_pool(
|
|
||||||
indexer_size,
|
|
||||||
c4_page_size,
|
|
||||||
dtype,
|
|
||||||
indexer_head_dim,
|
|
||||||
c4_layer_num,
|
|
||||||
device,
|
|
||||||
enable_memory_saver,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
self._init_compressed_layer_mapping()
|
self._init_compressed_layer_mapping()
|
||||||
@@ -804,53 +782,39 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool):
|
|||||||
data_lens: List[int] = []
|
data_lens: List[int] = []
|
||||||
item_lens: List[int] = []
|
item_lens: List[int] = []
|
||||||
|
|
||||||
if self._unified_kv:
|
def append_page_buffer(buf: torch.Tensor) -> None:
|
||||||
# Unified buffer per layer: [swa_pages + padded_compress_rows, head_dim].
|
assert buf.ndim == 2, f"expected 2D buffer, got {buf.ndim}D"
|
||||||
# Compressed region [swa_pages:] is page-contiguous (row swa_pages +
|
data_ptrs.append(buf.data_ptr())
|
||||||
# loc//ratio), so reuse the page-block PD transfer by offsetting the ptr
|
data_lens.append(buf.nbytes)
|
||||||
# past the SWA ring and setting item_len = one page of rows. The SWA ring
|
item_lens.append(buf[0].nbytes)
|
||||||
# ships separately as StateType.SWA_RING. Order [c4, c4_indexer, c128]
|
|
||||||
# mirrors the non-unified kv_data layout (keeps PP ptr-slicing valid).
|
|
||||||
stage_ratios = self.compression_ratios[self._stage_start : self._stage_end]
|
|
||||||
swa_pages = self.unified_kv_pool.swa_pages
|
|
||||||
|
|
||||||
def _append_compressed_entry(local_layer_id: int, ratio: int) -> None:
|
stage_ratios = self.compression_ratios[self._stage_start : self._stage_end]
|
||||||
buf = self.unified_kv_pool.kv_buffer[local_layer_id]
|
# Registration order defines the PD wire layout: C4 KV, C4 indexer, C128 KV.
|
||||||
assert buf.ndim == 2, f"expected 2D buffer, got {buf.ndim}D"
|
# Keep each indexer immediately after the KV buffers of the same ratio.
|
||||||
row_bytes = buf[0].nbytes
|
for ratio, kv_pool in self.kv_pools.items():
|
||||||
rows_per_page = self.page_size // ratio
|
if self._unified_kv:
|
||||||
compress_rows = buf.shape[0] - swa_pages
|
# Unified buffers store token rows after the SWA ring. Transfer
|
||||||
data_ptrs.append(buf.data_ptr() + swa_pages * row_bytes)
|
# compressed pages from the offset; SWA ships as StateType.SWA_RING.
|
||||||
data_lens.append(compress_rows * row_bytes)
|
swa_pages = self.unified_kv_pool.swa_pages
|
||||||
item_lens.append(rows_per_page * row_bytes)
|
for local_layer_id, layer_ratio in enumerate(stage_ratios):
|
||||||
|
if layer_ratio != ratio:
|
||||||
|
continue
|
||||||
|
buf = self.unified_kv_pool.kv_buffer[local_layer_id]
|
||||||
|
assert buf.ndim == 2, f"expected 2D buffer, got {buf.ndim}D"
|
||||||
|
row_bytes = buf[0].nbytes
|
||||||
|
rows_per_page = self.page_size // ratio
|
||||||
|
compress_rows = buf.shape[0] - swa_pages
|
||||||
|
data_ptrs.append(buf.data_ptr() + swa_pages * row_bytes)
|
||||||
|
data_lens.append(compress_rows * row_bytes)
|
||||||
|
item_lens.append(rows_per_page * row_bytes)
|
||||||
|
else:
|
||||||
|
for buf in kv_pool.kv_buffer:
|
||||||
|
append_page_buffer(buf)
|
||||||
|
|
||||||
c4_locals = [i for i, r in enumerate(stage_ratios) if r == 4]
|
indexer_pool = self.index_pools.get(ratio)
|
||||||
c128_locals = [i for i, r in enumerate(stage_ratios) if r == 128]
|
if indexer_pool is not None:
|
||||||
|
for buf in indexer_pool.contiguous_page_row_buffers():
|
||||||
for i in c4_locals:
|
append_page_buffer(buf)
|
||||||
_append_compressed_entry(i, 4)
|
|
||||||
for buf in self.c4_indexer_kv_pool.contiguous_page_row_buffers():
|
|
||||||
assert buf.ndim == 2, f"expected 2D buffer, got {buf.ndim}D"
|
|
||||||
data_ptrs.append(buf.data_ptr())
|
|
||||||
data_lens.append(buf.nbytes)
|
|
||||||
item_lens.append(buf[0].nbytes)
|
|
||||||
for i in c128_locals:
|
|
||||||
_append_compressed_entry(i, 128)
|
|
||||||
|
|
||||||
return data_ptrs, data_lens, item_lens
|
|
||||||
|
|
||||||
buf_groups = [
|
|
||||||
self.c4_kv_pool.kv_buffer,
|
|
||||||
self.c4_indexer_kv_pool.contiguous_page_row_buffers(),
|
|
||||||
self.c128_kv_pool.kv_buffer,
|
|
||||||
]
|
|
||||||
|
|
||||||
for bufs in buf_groups:
|
|
||||||
for buf in bufs:
|
|
||||||
assert buf.ndim == 2, f"expected 2D buffer, got {buf.ndim}D"
|
|
||||||
data_ptrs.append(buf.data_ptr())
|
|
||||||
data_lens.append(buf.nbytes)
|
|
||||||
item_lens.append(buf[0].nbytes)
|
|
||||||
|
|
||||||
return data_ptrs, data_lens, item_lens
|
return data_ptrs, data_lens, item_lens
|
||||||
|
|
||||||
@@ -950,6 +914,62 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool):
|
|||||||
item_lens.append(t[0].nbytes if ONLINE_C128 else t[0].nbytes * 128)
|
item_lens.append(t[0].nbytes if ONLINE_C128 else t[0].nbytes * 128)
|
||||||
return data_ptrs, data_lens, item_lens
|
return data_ptrs, data_lens, item_lens
|
||||||
|
|
||||||
|
def _init_compressed_pools(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
stage_ratios: Sequence[int],
|
||||||
|
page_size: int,
|
||||||
|
dtype: torch.dtype,
|
||||||
|
device: str,
|
||||||
|
enable_memory_saver: bool,
|
||||||
|
enable_hisparse: bool,
|
||||||
|
kv_pool_cls: type,
|
||||||
|
) -> None:
|
||||||
|
configs = self.compressed_pool_configs
|
||||||
|
layer_counts = {ratio: stage_ratios.count(ratio) for ratio in configs}
|
||||||
|
# Keep empty pools and allocation order for PP stages without a given ratio.
|
||||||
|
self.kv_pools: dict[int, Optional[DeepSeekV4SingleKVPool]] = {
|
||||||
|
ratio: None for ratio in configs
|
||||||
|
}
|
||||||
|
|
||||||
|
if not self._unified_kv:
|
||||||
|
for ratio, config in configs.items():
|
||||||
|
pool_cls = kv_pool_cls
|
||||||
|
if ratio == 4 and enable_hisparse:
|
||||||
|
assert not self.uniform_fp8, (
|
||||||
|
"enable_hisparse is not supported with --dsv4-attn-backend trtllm."
|
||||||
|
)
|
||||||
|
pool_cls = HiSparseC4DevicePool
|
||||||
|
self.kv_pools[ratio] = self._make_kv_pool(
|
||||||
|
size=config.kv_size,
|
||||||
|
page_size=page_size // ratio,
|
||||||
|
dtype=dtype,
|
||||||
|
layer_num=layer_counts[ratio],
|
||||||
|
device=device,
|
||||||
|
enable_memory_saver=enable_memory_saver,
|
||||||
|
global_page_size=page_size,
|
||||||
|
cls=pool_cls,
|
||||||
|
)
|
||||||
|
|
||||||
|
self.index_pools: dict[int, DeepSeekV4IndexerPool] = {
|
||||||
|
ratio: self._make_indexer_pool(
|
||||||
|
config.indexer_size,
|
||||||
|
page_size // ratio,
|
||||||
|
dtype,
|
||||||
|
self.indexer_head_dim,
|
||||||
|
layer_counts[ratio],
|
||||||
|
device,
|
||||||
|
enable_memory_saver,
|
||||||
|
)
|
||||||
|
for ratio, config in configs.items()
|
||||||
|
if config.indexer_size is not None
|
||||||
|
}
|
||||||
|
|
||||||
|
# HiCache and hardware backends still access the per-ratio attributes.
|
||||||
|
self.c4_kv_pool = self.kv_pools[4]
|
||||||
|
self.c128_kv_pool = self.kv_pools[128]
|
||||||
|
self.c4_indexer_kv_pool = self.index_pools[4]
|
||||||
|
|
||||||
def _make_kv_pool(
|
def _make_kv_pool(
|
||||||
self,
|
self,
|
||||||
*,
|
*,
|
||||||
@@ -1002,21 +1022,17 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool):
|
|||||||
enable_memory_saver,
|
enable_memory_saver,
|
||||||
)
|
)
|
||||||
|
|
||||||
def _state_pool_size(self, ratio: int) -> int:
|
def _make_compress_state_pool(
|
||||||
return self.c4_state_pool_size if ratio == 4 else self.c128_state_pool_size
|
self, ratio: int, *, head_dim: int, enable_memory_saver: bool
|
||||||
|
|
||||||
def _make_attn_state_pool(
|
|
||||||
self, ratio: int, enable_memory_saver: bool
|
|
||||||
) -> CompressStatePool:
|
) -> CompressStatePool:
|
||||||
"""Build the per-layer attention compress-state pool for ``ratio``
|
"""Build attention or indexer state; hardware backends override this factory."""
|
||||||
(4 or 128). Overridden by :class:`DSV4NPUTokenToKVPool` to swap the
|
config = self.compressed_pool_configs[ratio]
|
||||||
ring-buffered pool for the NPU paged one."""
|
|
||||||
return CompressStatePool(
|
return CompressStatePool(
|
||||||
size=self._state_pool_size(ratio),
|
size=config.state_size,
|
||||||
ring_size=self.get_ring_size(ratio),
|
ring_size=self.get_ring_size(ratio),
|
||||||
overlap=ratio == 4,
|
overlap=ratio == 4,
|
||||||
head_dim=self.qk_nope_head_dim + self.qk_rope_head_dim,
|
head_dim=head_dim,
|
||||||
dtype=self.c4_state_dtype if ratio == 4 else self.c128_state_dtype,
|
dtype=config.state_dtype,
|
||||||
device=self.device,
|
device=self.device,
|
||||||
enable_memory_saver=enable_memory_saver,
|
enable_memory_saver=enable_memory_saver,
|
||||||
ratio=ratio,
|
ratio=ratio,
|
||||||
@@ -1027,25 +1043,7 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool):
|
|||||||
),
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
def _make_indexer_state_pool(
|
|
||||||
self, ratio: int, enable_memory_saver: bool
|
|
||||||
) -> CompressStatePool:
|
|
||||||
"""Build the per-layer indexer compress-state pool (c4 only)."""
|
|
||||||
return CompressStatePool(
|
|
||||||
size=self._state_pool_size(ratio),
|
|
||||||
ring_size=self.get_ring_size(ratio),
|
|
||||||
overlap=ratio == 4,
|
|
||||||
head_dim=self.indexer_head_dim,
|
|
||||||
device=self.device,
|
|
||||||
dtype=self.c4_state_dtype,
|
|
||||||
enable_memory_saver=enable_memory_saver,
|
|
||||||
ratio=ratio,
|
|
||||||
swa_page_size=self.swa_page_size,
|
|
||||||
)
|
|
||||||
|
|
||||||
def _init_paged_compress_states(self, enable_memory_saver: bool):
|
def _init_paged_compress_states(self, enable_memory_saver: bool):
|
||||||
c4_state_pool_size = self.c4_state_pool_size
|
|
||||||
c128_state_pool_size = self.c128_state_pool_size
|
|
||||||
total_L = len(self.compression_ratios)
|
total_L = len(self.compression_ratios)
|
||||||
self.compress_state_pools: List[Optional[CompressStatePool]] = [None] * total_L
|
self.compress_state_pools: List[Optional[CompressStatePool]] = [None] * total_L
|
||||||
self.indexer_compress_state_pools: List[Optional[CompressStatePool]] = [
|
self.indexer_compress_state_pools: List[Optional[CompressStatePool]] = [
|
||||||
@@ -1057,44 +1055,34 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool):
|
|||||||
if ratio == 0:
|
if ratio == 0:
|
||||||
continue
|
continue
|
||||||
|
|
||||||
self.compress_state_pools[idx] = self._make_attn_state_pool(
|
self.compress_state_pools[idx] = self._make_compress_state_pool(
|
||||||
ratio, enable_memory_saver
|
ratio,
|
||||||
|
head_dim=self.qk_nope_head_dim + self.qk_rope_head_dim,
|
||||||
|
enable_memory_saver=enable_memory_saver,
|
||||||
)
|
)
|
||||||
|
|
||||||
if ratio == 4:
|
if ratio in self.index_pools:
|
||||||
self.indexer_compress_state_pools[idx] = self._make_indexer_state_pool(
|
self.indexer_compress_state_pools[idx] = self._make_compress_state_pool(
|
||||||
ratio, enable_memory_saver
|
ratio,
|
||||||
|
head_dim=self.indexer_head_dim,
|
||||||
|
enable_memory_saver=enable_memory_saver,
|
||||||
)
|
)
|
||||||
|
|
||||||
def _init_compressed_layer_mapping(self):
|
def _init_compressed_layer_mapping(self):
|
||||||
c0_cnt = c4_cnt = c128_cnt = 0
|
layer_counts = {0: 0, **{ratio: 0 for ratio in self.kv_pools}}
|
||||||
total_L = len(self.compression_ratios)
|
total_L = len(self.compression_ratios)
|
||||||
self.layer_mapping: List[Optional[DeepSeekV4LayerItem]] = [None] * total_L
|
self.layer_mapping: List[Optional[DeepSeekV4LayerItem]] = [None] * total_L
|
||||||
|
|
||||||
for idx in range(self._stage_start, self._stage_end):
|
for idx in range(self._stage_start, self._stage_end):
|
||||||
ratio = self.compression_ratios[idx]
|
ratio = self.compression_ratios[idx]
|
||||||
if ratio == 0:
|
if ratio not in layer_counts:
|
||||||
self.layer_mapping[idx] = DeepSeekV4LayerItem(
|
|
||||||
compress_ratio=0,
|
|
||||||
compress_layer_id=c0_cnt,
|
|
||||||
)
|
|
||||||
c0_cnt += 1
|
|
||||||
elif ratio == 4:
|
|
||||||
self.layer_mapping[idx] = DeepSeekV4LayerItem(
|
|
||||||
compress_ratio=4,
|
|
||||||
compress_layer_id=c4_cnt,
|
|
||||||
compress_kv_pool=self.c4_kv_pool,
|
|
||||||
)
|
|
||||||
c4_cnt += 1
|
|
||||||
elif ratio == 128:
|
|
||||||
self.layer_mapping[idx] = DeepSeekV4LayerItem(
|
|
||||||
compress_ratio=128,
|
|
||||||
compress_layer_id=c128_cnt,
|
|
||||||
compress_kv_pool=self.c128_kv_pool,
|
|
||||||
)
|
|
||||||
c128_cnt += 1
|
|
||||||
else:
|
|
||||||
raise ValueError(f"Unsupported compression ratio: {ratio}")
|
raise ValueError(f"Unsupported compression ratio: {ratio}")
|
||||||
|
self.layer_mapping[idx] = DeepSeekV4LayerItem(
|
||||||
|
compress_ratio=ratio,
|
||||||
|
compress_layer_id=layer_counts[ratio],
|
||||||
|
compress_kv_pool=self.kv_pools.get(ratio),
|
||||||
|
)
|
||||||
|
layer_counts[ratio] += 1
|
||||||
|
|
||||||
def wait_layer_transfer(self, layer_id: int) -> None:
|
def wait_layer_transfer(self, layer_id: int) -> None:
|
||||||
if self.layer_transfer_counter is not None:
|
if self.layer_transfer_counter is not None:
|
||||||
@@ -1213,20 +1201,6 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool):
|
|||||||
def get_swa_raw_buffer(self, layer_id: int) -> torch.Tensor:
|
def get_swa_raw_buffer(self, layer_id: int) -> torch.Tensor:
|
||||||
return self.swa_kv_pool.kv_buffer[self._swa_local_layer_id(layer_id)]
|
return self.swa_kv_pool.kv_buffer[self._swa_local_layer_id(layer_id)]
|
||||||
|
|
||||||
def get_swa_key_buffer(self, layer_id: int) -> torch.Tensor:
|
|
||||||
self.wait_layer_transfer(layer_id)
|
|
||||||
return self.swa_kv_pool.get_key_buffer(self._swa_local_layer_id(layer_id))
|
|
||||||
|
|
||||||
def set_swa_key_buffer(
|
|
||||||
self,
|
|
||||||
layer_id: int,
|
|
||||||
loc: torch.Tensor,
|
|
||||||
cache_nope_fp8_rope_bf16_pack: NopeFp8RopeBf16Pack,
|
|
||||||
) -> None:
|
|
||||||
self.swa_kv_pool.set_key_buffer(
|
|
||||||
self._swa_local_layer_id(layer_id), loc, cache_nope_fp8_rope_bf16_pack
|
|
||||||
)
|
|
||||||
|
|
||||||
def get_extra_key_page_size(self, layer_id: int) -> int:
|
def get_extra_key_page_size(self, layer_id: int) -> int:
|
||||||
_, _, compress_kv_pool = self.layer_mapping[layer_id]
|
_, _, compress_kv_pool = self.layer_mapping[layer_id]
|
||||||
assert compress_kv_pool is not None
|
assert compress_kv_pool is not None
|
||||||
@@ -1250,26 +1224,36 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool):
|
|||||||
compress_layer_id, loc, cache_nope_fp8_rope_bf16_pack
|
compress_layer_id, loc, cache_nope_fp8_rope_bf16_pack
|
||||||
)
|
)
|
||||||
|
|
||||||
def get_index_k_page_size(self) -> int:
|
def _indexer_pool(self, compress_ratio: int) -> DeepSeekV4IndexerPool:
|
||||||
return self.c4_indexer_kv_pool.page_size
|
pool = self.index_pools.get(compress_ratio)
|
||||||
|
assert pool is not None, (
|
||||||
|
f"No indexer pool for compression ratio {compress_ratio}"
|
||||||
|
)
|
||||||
|
return pool
|
||||||
|
|
||||||
|
def get_index_k_page_size(self, compress_ratio: int = 4) -> int:
|
||||||
|
return self._indexer_pool(compress_ratio).page_size
|
||||||
|
|
||||||
def get_index_k_with_scale_buffer(self, layer_id: int) -> torch.Tensor:
|
def get_index_k_with_scale_buffer(self, layer_id: int) -> torch.Tensor:
|
||||||
self.wait_layer_transfer(layer_id)
|
self.wait_layer_transfer(layer_id)
|
||||||
compress_ratio, compress_layer_id, _ = self.layer_mapping[layer_id]
|
compress_ratio, compress_layer_id, _ = self.layer_mapping[layer_id]
|
||||||
assert compress_ratio == 4, f"only c4 has indexer, got {compress_ratio = }"
|
return self._indexer_pool(compress_ratio).get_index_k_with_scale_buffer(
|
||||||
return self.c4_indexer_kv_pool.get_index_k_with_scale_buffer(compress_layer_id)
|
compress_layer_id
|
||||||
|
)
|
||||||
|
|
||||||
def get_index_k_fp4_payload_buffer(self, layer_id: int) -> torch.Tensor:
|
def get_index_k_fp4_payload_buffer(self, layer_id: int) -> torch.Tensor:
|
||||||
self.wait_layer_transfer(layer_id)
|
self.wait_layer_transfer(layer_id)
|
||||||
compress_ratio, compress_layer_id, _ = self.layer_mapping[layer_id]
|
compress_ratio, compress_layer_id, _ = self.layer_mapping[layer_id]
|
||||||
assert compress_ratio == 4, f"only c4 has indexer, got {compress_ratio = }"
|
return self._indexer_pool(compress_ratio).get_index_k_fp4_payload_buffer(
|
||||||
return self.c4_indexer_kv_pool.get_index_k_fp4_payload_buffer(compress_layer_id)
|
compress_layer_id
|
||||||
|
)
|
||||||
|
|
||||||
def get_index_k_fp4_scale_buffer(self, layer_id: int) -> torch.Tensor:
|
def get_index_k_fp4_scale_buffer(self, layer_id: int) -> torch.Tensor:
|
||||||
self.wait_layer_transfer(layer_id)
|
self.wait_layer_transfer(layer_id)
|
||||||
compress_ratio, compress_layer_id, _ = self.layer_mapping[layer_id]
|
compress_ratio, compress_layer_id, _ = self.layer_mapping[layer_id]
|
||||||
assert compress_ratio == 4, f"only c4 has indexer, got {compress_ratio = }"
|
return self._indexer_pool(compress_ratio).get_index_k_fp4_scale_buffer(
|
||||||
return self.c4_indexer_kv_pool.get_index_k_fp4_scale_buffer(compress_layer_id)
|
compress_layer_id
|
||||||
|
)
|
||||||
|
|
||||||
def get_index_k_scale_buffer(
|
def get_index_k_scale_buffer(
|
||||||
self,
|
self,
|
||||||
@@ -1281,8 +1265,7 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool):
|
|||||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||||
self.wait_layer_transfer(layer_id)
|
self.wait_layer_transfer(layer_id)
|
||||||
compress_ratio, compress_layer_id, _ = self.layer_mapping[layer_id]
|
compress_ratio, compress_layer_id, _ = self.layer_mapping[layer_id]
|
||||||
assert compress_ratio == 4, f"only c4 has indexer, got {compress_ratio = }"
|
return self._indexer_pool(compress_ratio).get_index_k_scale_buffer(
|
||||||
return self.c4_indexer_kv_pool.get_index_k_scale_buffer(
|
|
||||||
compress_layer_id,
|
compress_layer_id,
|
||||||
seq_len_tensor,
|
seq_len_tensor,
|
||||||
page_indices,
|
page_indices,
|
||||||
@@ -1298,8 +1281,7 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool):
|
|||||||
index_k_scale: torch.Tensor,
|
index_k_scale: torch.Tensor,
|
||||||
) -> None:
|
) -> None:
|
||||||
compress_ratio, compress_layer_id, _ = self.layer_mapping[layer_id]
|
compress_ratio, compress_layer_id, _ = self.layer_mapping[layer_id]
|
||||||
assert compress_ratio == 4, f"only c4 has indexer, got {compress_ratio = }"
|
self._indexer_pool(compress_ratio).set_index_k_scale_buffer(
|
||||||
self.c4_indexer_kv_pool.set_index_k_scale_buffer(
|
|
||||||
compress_layer_id, loc, index_k, index_k_scale
|
compress_layer_id, loc, index_k, index_k_scale
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -1425,8 +1407,9 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool):
|
|||||||
cache_k: torch.Tensor,
|
cache_k: torch.Tensor,
|
||||||
) -> None:
|
) -> None:
|
||||||
compress_ratio, compress_layer_id, _ = self.layer_mapping[layer_id]
|
compress_ratio, compress_layer_id, _ = self.layer_mapping[layer_id]
|
||||||
assert compress_ratio == 4, f"only c4 has indexer, got {compress_ratio = }"
|
return self._indexer_pool(compress_ratio).set_index_fused(
|
||||||
return self.c4_indexer_kv_pool.set_index_fused(compress_layer_id, loc, cache_k)
|
compress_layer_id, loc, cache_k
|
||||||
|
)
|
||||||
|
|
||||||
def set_index_k_fp4(
|
def set_index_k_fp4(
|
||||||
self,
|
self,
|
||||||
@@ -1435,5 +1418,6 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool):
|
|||||||
cache_k: torch.Tensor,
|
cache_k: torch.Tensor,
|
||||||
) -> None:
|
) -> None:
|
||||||
compress_ratio, compress_layer_id, _ = self.layer_mapping[layer_id]
|
compress_ratio, compress_layer_id, _ = self.layer_mapping[layer_id]
|
||||||
assert compress_ratio == 4, f"only c4 has indexer, got {compress_ratio = }"
|
return self._indexer_pool(compress_ratio).set_index_fp4(
|
||||||
return self.c4_indexer_kv_pool.set_index_fp4(compress_layer_id, loc, cache_k)
|
compress_layer_id, loc, cache_k
|
||||||
|
)
|
||||||
|
|||||||
@@ -45,6 +45,7 @@ class TestDSV4PagedIndexerMetadata(CustomTestCase):
|
|||||||
):
|
):
|
||||||
metadata = PagedIndexerMetadata(
|
metadata = PagedIndexerMetadata(
|
||||||
page_size=256,
|
page_size=256,
|
||||||
|
compressed_page_size=64,
|
||||||
page_table=torch.zeros((1, 1), dtype=torch.int32),
|
page_table=torch.zeros((1, 1), dtype=torch.int32),
|
||||||
compressed_seq_lens=torch.tensor([65], dtype=torch.int32),
|
compressed_seq_lens=torch.tensor([65], dtype=torch.int32),
|
||||||
use_topk_v2=False,
|
use_topk_v2=False,
|
||||||
@@ -59,22 +60,7 @@ class TestDSV4PagedIndexerMetadata(CustomTestCase):
|
|||||||
self.assertEqual(args[1:], (64, 1))
|
self.assertEqual(args[1:], (64, 1))
|
||||||
jit_metadata.assert_not_called()
|
jit_metadata.assert_not_called()
|
||||||
|
|
||||||
def test_sm120_fp8_torch_fallback_keeps_metadata_none(self):
|
def test_torch_fallback_skips_deep_gemm_and_ineligible_topk_plan(self):
|
||||||
with (
|
|
||||||
envs.SGLANG_FP8_PAGED_MQA_LOGITS_TORCH.override(True),
|
|
||||||
envs.SGLANG_OPT_USE_AITER_INDEXER.override(False),
|
|
||||||
envs.SGLANG_OPT_USE_TOPK_V2.override(False),
|
|
||||||
):
|
|
||||||
metadata = PagedIndexerMetadata(
|
|
||||||
page_size=256,
|
|
||||||
page_table=torch.zeros((1, 1), dtype=torch.int32),
|
|
||||||
compressed_seq_lens=torch.tensor([65], dtype=torch.int32),
|
|
||||||
use_topk_v2=False,
|
|
||||||
)
|
|
||||||
|
|
||||||
self.assertIsNone(metadata.deep_gemm_metadata)
|
|
||||||
|
|
||||||
def test_topk_v2_ineligible_backend_skips_plan(self):
|
|
||||||
with (
|
with (
|
||||||
envs.SGLANG_FP8_PAGED_MQA_LOGITS_TORCH.override(True),
|
envs.SGLANG_FP8_PAGED_MQA_LOGITS_TORCH.override(True),
|
||||||
envs.SGLANG_OPT_USE_AITER_INDEXER.override(False),
|
envs.SGLANG_OPT_USE_AITER_INDEXER.override(False),
|
||||||
@@ -83,14 +69,54 @@ class TestDSV4PagedIndexerMetadata(CustomTestCase):
|
|||||||
):
|
):
|
||||||
metadata = PagedIndexerMetadata(
|
metadata = PagedIndexerMetadata(
|
||||||
page_size=256,
|
page_size=256,
|
||||||
|
compressed_page_size=64,
|
||||||
page_table=torch.zeros((1, 1), dtype=torch.int32),
|
page_table=torch.zeros((1, 1), dtype=torch.int32),
|
||||||
compressed_seq_lens=torch.tensor([65], dtype=torch.int32),
|
compressed_seq_lens=torch.tensor([65], dtype=torch.int32),
|
||||||
use_topk_v2=False,
|
use_topk_v2=False,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
self.assertIsNone(metadata.deep_gemm_metadata)
|
||||||
plan_topk_v2.assert_not_called()
|
plan_topk_v2.assert_not_called()
|
||||||
self.assertEqual(metadata.topk_metadata.numel(), 0)
|
self.assertEqual(metadata.topk_metadata.numel(), 0)
|
||||||
|
|
||||||
|
def test_physical_page_size_controls_metadata_and_replay(self):
|
||||||
|
planner = MagicMock(return_value=torch.zeros((1, 2), dtype=torch.int32))
|
||||||
|
deep_gemm = SimpleNamespace(
|
||||||
|
get_num_sms=MagicMock(return_value=1),
|
||||||
|
get_paged_mqa_logits_metadata=planner,
|
||||||
|
)
|
||||||
|
with patch.dict(sys.modules, {"deep_gemm": deep_gemm}):
|
||||||
|
metadata = [
|
||||||
|
PagedIndexerMetadata(
|
||||||
|
page_size=256,
|
||||||
|
compressed_page_size=page_size,
|
||||||
|
page_table=torch.zeros((1, 3), dtype=torch.int32),
|
||||||
|
compressed_seq_lens=torch.tensor([65], dtype=torch.int32),
|
||||||
|
use_topk_v2=False,
|
||||||
|
force_deep_gemm_metadata=True,
|
||||||
|
)
|
||||||
|
for page_size in (64, 32, 32)
|
||||||
|
]
|
||||||
|
|
||||||
|
self.assertEqual(
|
||||||
|
[call.args[1] for call in planner.call_args_list], [64, 32, 32]
|
||||||
|
)
|
||||||
|
self.assertEqual([m.max_compressed_seq_len for m in metadata], [192, 96, 96])
|
||||||
|
self.assertEqual([m.max_seq_len for m in metadata], [768, 768, 768])
|
||||||
|
with self.assertRaisesRegex(AssertionError, "compressed_page_size"):
|
||||||
|
metadata[0].copy_(metadata[1])
|
||||||
|
|
||||||
|
destination, source = metadata[1:]
|
||||||
|
source.page_table.fill_(7)
|
||||||
|
source.compressed_seq_lens.fill_(17)
|
||||||
|
page_table_ptr = destination.page_table.data_ptr()
|
||||||
|
destination.copy_(source)
|
||||||
|
self.assertEqual(destination.page_table.data_ptr(), page_table_ptr)
|
||||||
|
torch.testing.assert_close(destination.page_table, source.page_table)
|
||||||
|
torch.testing.assert_close(
|
||||||
|
destination.compressed_seq_lens, source.compressed_seq_lens
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
class TestDSV4FlashInferTopK(CustomTestCase):
|
class TestDSV4FlashInferTopK(CustomTestCase):
|
||||||
def test_compact_page_transform_respects_fuse_topk(self):
|
def test_compact_page_transform_respects_fuse_topk(self):
|
||||||
|
|||||||
@@ -0,0 +1,180 @@
|
|||||||
|
import unittest
|
||||||
|
from itertools import product
|
||||||
|
from types import SimpleNamespace
|
||||||
|
from unittest.mock import MagicMock, patch
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from sglang.srt.mem_cache.deepseek_v4_memory_pool import (
|
||||||
|
DeepSeekV4SingleKVPool,
|
||||||
|
DeepSeekV4TokenToKVPool,
|
||||||
|
_CompressedPoolConfig,
|
||||||
|
)
|
||||||
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
|
from sglang.test.test_utils import CustomTestCase
|
||||||
|
|
||||||
|
register_cpu_ci(est_time=5, suite="base-a-test-cpu")
|
||||||
|
|
||||||
|
|
||||||
|
class TestDSV4CompressedPools(CustomTestCase):
|
||||||
|
def test_pp_mapping_and_pd_buffer_order(self):
|
||||||
|
for unified, stage_ratios in product(
|
||||||
|
(False, True), ([4, 0, 128, 4], [128], [0])
|
||||||
|
):
|
||||||
|
with self.subTest(unified=unified, stage_ratios=stage_ratios):
|
||||||
|
pool = DeepSeekV4TokenToKVPool.__new__(DeepSeekV4TokenToKVPool)
|
||||||
|
pool._unified_kv = unified
|
||||||
|
pool.uniform_fp8 = False
|
||||||
|
pool.compressed_pool_configs = {
|
||||||
|
4: _CompressedPoolConfig(
|
||||||
|
256, 64, torch.bfloat16, indexer_size=1024
|
||||||
|
),
|
||||||
|
128: _CompressedPoolConfig(512, 8, torch.float32),
|
||||||
|
}
|
||||||
|
pool.indexer_head_dim = 128
|
||||||
|
pool.page_size = 256
|
||||||
|
pool.compression_ratios = [128] + stage_ratios + [4]
|
||||||
|
pool._stage_start = 1
|
||||||
|
pool._stage_end = 1 + len(stage_ratios)
|
||||||
|
|
||||||
|
def kv_factory(**kwargs):
|
||||||
|
return SimpleNamespace(
|
||||||
|
kv_buffer=[
|
||||||
|
torch.empty((3, kwargs["page_size"]), dtype=torch.uint8)
|
||||||
|
for _ in range(kwargs["layer_num"])
|
||||||
|
]
|
||||||
|
)
|
||||||
|
|
||||||
|
# Separate payload/scale buffers exercise the FP4 transfer contract.
|
||||||
|
indexer_buffers = [
|
||||||
|
torch.empty((3, width), dtype=torch.uint8)
|
||||||
|
for _ in range(stage_ratios.count(4))
|
||||||
|
for width in (32, 2)
|
||||||
|
]
|
||||||
|
indexer = SimpleNamespace(
|
||||||
|
contiguous_page_row_buffers=lambda: indexer_buffers
|
||||||
|
)
|
||||||
|
with (
|
||||||
|
patch.object(pool, "_make_kv_pool", side_effect=kv_factory),
|
||||||
|
patch.object(pool, "_make_indexer_pool", return_value=indexer),
|
||||||
|
):
|
||||||
|
pool._init_compressed_pools(
|
||||||
|
stage_ratios=stage_ratios,
|
||||||
|
page_size=256,
|
||||||
|
dtype=torch.float8_e4m3fn,
|
||||||
|
device="cpu",
|
||||||
|
enable_memory_saver=False,
|
||||||
|
enable_hisparse=False,
|
||||||
|
kv_pool_cls=DeepSeekV4SingleKVPool,
|
||||||
|
)
|
||||||
|
pool._init_compressed_layer_mapping()
|
||||||
|
self.assertIsNone(pool.layer_mapping[0])
|
||||||
|
self.assertIsNone(pool.layer_mapping[-1])
|
||||||
|
for local_id, ratio in enumerate(stage_ratios):
|
||||||
|
item = pool.layer_mapping[local_id + 1]
|
||||||
|
self.assertEqual(
|
||||||
|
item.compress_layer_id, stage_ratios[:local_id].count(ratio)
|
||||||
|
)
|
||||||
|
self.assertIs(item.compress_kv_pool, pool.kv_pools.get(ratio))
|
||||||
|
self.assertIs(pool.c4_kv_pool, pool.kv_pools[4])
|
||||||
|
self.assertIs(pool.c128_kv_pool, pool.kv_pools[128])
|
||||||
|
self.assertIs(pool.c4_indexer_kv_pool, pool.index_pools[4])
|
||||||
|
|
||||||
|
if unified:
|
||||||
|
buffers = [
|
||||||
|
torch.empty((9, 8), dtype=torch.uint8) for _ in stage_ratios
|
||||||
|
]
|
||||||
|
pool.unified_kv_pool = SimpleNamespace(
|
||||||
|
swa_pages=2, kv_buffer=buffers
|
||||||
|
)
|
||||||
|
|
||||||
|
def kv_entries(ratio):
|
||||||
|
return [
|
||||||
|
(buf.data_ptr() + 16, 56, 256 // ratio * 8)
|
||||||
|
for buf, r in zip(buffers, stage_ratios)
|
||||||
|
if r == ratio
|
||||||
|
]
|
||||||
|
else:
|
||||||
|
|
||||||
|
def kv_entries(ratio):
|
||||||
|
return [
|
||||||
|
(b.data_ptr(), b.nbytes, b[0].nbytes)
|
||||||
|
for b in pool.kv_pools[ratio].kv_buffer
|
||||||
|
]
|
||||||
|
|
||||||
|
indexer_entries = [
|
||||||
|
(b.data_ptr(), b.nbytes, b[0].nbytes) for b in indexer_buffers
|
||||||
|
]
|
||||||
|
expected = kv_entries(4) + indexer_entries + kv_entries(128)
|
||||||
|
actual = list(zip(*pool.get_contiguous_buf_infos()))
|
||||||
|
self.assertEqual(actual, expected)
|
||||||
|
|
||||||
|
def test_shared_state_factory_preserves_layouts(self):
|
||||||
|
pool = DeepSeekV4TokenToKVPool.__new__(DeepSeekV4TokenToKVPool)
|
||||||
|
pool.compressed_pool_configs = {
|
||||||
|
4: _CompressedPoolConfig(256, 64, torch.bfloat16, indexer_size=1024),
|
||||||
|
128: _CompressedPoolConfig(512, 8, torch.float32),
|
||||||
|
}
|
||||||
|
pool.compression_ratios = [0, 4, 128]
|
||||||
|
pool._stage_start, pool._stage_end = 0, 3
|
||||||
|
pool.index_pools = {4: object()}
|
||||||
|
pool.qk_nope_head_dim, pool.qk_rope_head_dim = 448, 64
|
||||||
|
pool.indexer_head_dim = 128
|
||||||
|
pool.device = "cpu"
|
||||||
|
pool.swa_page_size = 128
|
||||||
|
pool.online_mtp_max_draft_tokens = 3
|
||||||
|
for online in (False, True):
|
||||||
|
with (
|
||||||
|
self.subTest(online=online),
|
||||||
|
patch(
|
||||||
|
"sglang.srt.mem_cache.deepseek_v4_memory_pool.ONLINE_C128", online
|
||||||
|
),
|
||||||
|
patch.object(
|
||||||
|
pool,
|
||||||
|
"get_ring_size",
|
||||||
|
side_effect=lambda r: 8 if r == 4 else (1 if online else 128),
|
||||||
|
),
|
||||||
|
):
|
||||||
|
pool._init_paged_compress_states(False)
|
||||||
|
c4 = pool.compress_state_pools[1].kv_score_buffer.kv_score
|
||||||
|
indexer = pool.indexer_compress_state_pools[1].kv_score_buffer.kv_score
|
||||||
|
c128 = pool.compress_state_pools[2].kv_score_buffer.kv_score
|
||||||
|
self.assertEqual(c4.shape, (76, 2048))
|
||||||
|
self.assertEqual(indexer.shape, (76, 512))
|
||||||
|
self.assertEqual(c128.shape, (40, 1536) if online else (256, 1024))
|
||||||
|
self.assertEqual(c4.dtype, torch.bfloat16)
|
||||||
|
self.assertEqual(indexer.dtype, torch.bfloat16)
|
||||||
|
self.assertEqual(c128.dtype, torch.float32)
|
||||||
|
self.assertNotEqual(c4.data_ptr(), indexer.data_ptr())
|
||||||
|
self.assertIsNone(pool.compress_state_pools[0])
|
||||||
|
self.assertIsNone(pool.indexer_compress_state_pools[2])
|
||||||
|
|
||||||
|
def test_indexer_access_uses_layer_ratio_and_waits_only_for_reads(self):
|
||||||
|
pool = DeepSeekV4TokenToKVPool.__new__(DeepSeekV4TokenToKVPool)
|
||||||
|
pool.kv_pools = {4: None, 128: None}
|
||||||
|
pool.compression_ratios = [4, 128, 4]
|
||||||
|
pool._stage_start, pool._stage_end = 0, 3
|
||||||
|
pool._init_compressed_layer_mapping()
|
||||||
|
indexer = MagicMock(page_size=32)
|
||||||
|
pool.index_pools = {4: indexer}
|
||||||
|
trace = MagicMock()
|
||||||
|
trace.attach_mock(indexer, "indexer")
|
||||||
|
with patch.object(pool, "wait_layer_transfer") as wait:
|
||||||
|
trace.attach_mock(wait, "wait")
|
||||||
|
pool.get_index_k_fp4_payload_buffer(2)
|
||||||
|
pool.set_index_k_fp4(2, "loc", "cache")
|
||||||
|
self.assertEqual(
|
||||||
|
trace.mock_calls,
|
||||||
|
[
|
||||||
|
unittest.mock.call.wait(2),
|
||||||
|
unittest.mock.call.indexer.get_index_k_fp4_payload_buffer(1),
|
||||||
|
unittest.mock.call.indexer.set_index_fp4(1, "loc", "cache"),
|
||||||
|
],
|
||||||
|
)
|
||||||
|
self.assertEqual(pool.get_index_k_page_size(4), 32)
|
||||||
|
with self.assertRaisesRegex(AssertionError, "No indexer pool"):
|
||||||
|
pool.get_index_k_page_size(128)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
Reference in New Issue
Block a user