Facade DSA index-cache: MTP topk-reuse state + index-K storage (#28609)

This commit is contained in:
Xinyuan Tong
2026-08-06 00:34:31 -07:00
committed by GitHub
parent 735995e7bd
commit 31c1e5943f
9 changed files with 672 additions and 388 deletions
@@ -0,0 +1,86 @@
from __future__ import annotations
from contextlib import contextmanager
from typing import TYPE_CHECKING, Iterator, Optional
if TYPE_CHECKING:
import torch
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
class IndexTopKShareState:
def __init__(
self,
forward_batch: ForwardBatch,
topk_indices: Optional[torch.Tensor],
):
self._forward_batch = forward_batch
self._topk_indices = topk_indices
@classmethod
def from_mtp_carry(cls, forward_batch: ForwardBatch) -> IndexTopKShareState:
topk_indices = (
forward_batch.spec_info.dsa_topk_indices
if forward_batch.reuse_dsa_topk_indices
else None
)
return cls(forward_batch, topk_indices)
@property
def enabled(self) -> bool:
return bool(self._forward_batch.reuse_dsa_topk_indices)
@property
def _seed_buf(self) -> Optional[torch.Tensor]:
if self._forward_batch.forward_mode.is_extend(include_draft_extend_v2=True):
return self._forward_batch.spec_info.dsa_seed_topk_capture
return None
@property
def should_publish(self) -> bool:
return self.enabled or self._seed_buf is not None
@property
def topk_indices(self) -> Optional[torch.Tensor]:
return self._topk_indices
def update(self, topk_indices: Optional[torch.Tensor]) -> None:
self._topk_indices = topk_indices
def publish(self) -> None:
if self._topk_indices is None or not self.should_publish:
return
if self.enabled:
self._forward_batch.spec_info.dsa_topk_indices = self._topk_indices
seed_buf = self._seed_buf
if seed_buf is not None:
sel = self._forward_batch.spec_info.dsa_seed_topk_select
src = (
self._topk_indices[: seed_buf.shape[0]]
if sel is None
else self._topk_indices[sel]
)
seed_buf[: src.shape[0]].copy_(src)
@classmethod
@contextmanager
def mtp_iteration(
cls,
forward_batch: ForwardBatch,
enabled: bool = True,
keep_carry_seed: bool = False,
) -> Iterator[Optional[IndexTopKShareState]]:
if not enabled:
yield None
return
spec_info = forward_batch.spec_info
forward_batch.reuse_dsa_topk_indices = True
# Keep the draft-extend seed so step 0 reuses it; else recompute it.
if not (keep_carry_seed and spec_info.dsa_topk_indices is not None):
spec_info.dsa_topk_indices = None
try:
yield cls.from_mtp_carry(forward_batch)
finally:
spec_info.dsa_topk_indices = None
forward_batch.reuse_dsa_topk_indices = False
@@ -37,6 +37,7 @@ import torch
from sglang.kernels.ops.attention.dsa import index_buf_accessor from sglang.kernels.ops.attention.dsa import index_buf_accessor
from sglang.srt.layers.cp.utils import get_layer_owner, get_layer_shard_range from sglang.srt.layers.cp.utils import get_layer_owner, get_layer_shard_range
from sglang.srt.mem_cache.index_key_cache import IndexKeyCache
from sglang.srt.mem_cache.memory_pool import ( from sglang.srt.mem_cache.memory_pool import (
GPU_MEMORY_TYPE_KV_CACHE, GPU_MEMORY_TYPE_KV_CACHE,
DSATokenToKVPool, DSATokenToKVPool,
@@ -53,6 +54,141 @@ if TYPE_CHECKING:
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
class LayerSplitIndexKeyCache(IndexKeyCache):
def __init__(self, pool: LayerSplitDSATokenToKVPool, index_buf_size: int):
super().__init__(pool, index_buf_size)
num_pages = (index_buf_size + pool.page_size + 1) // pool.page_size
with (
torch.cuda.use_mem_pool(pool.custom_mem_pool)
if pool.custom_mem_pool
else nullcontext()
):
self.remote_buffer = torch.empty(
self._buffer_shape(num_pages),
dtype=pool.index_k_with_scale_buffer_dtype,
device=pool.device,
)
self.remote_layer_id: Optional[int] = None
def _layer_num_pages(self, layer_idx: int, num_pages: int) -> int:
layer_id = self.pool.start_layer + layer_idx
return num_pages if self.pool._is_layer_owned(layer_id) else 0
def clear(self) -> None:
super().clear()
del self.remote_buffer
def move(self, tgt_loc: torch.Tensor, src_loc: torch.Tensor) -> None:
if tgt_loc.numel() == 0:
return
tgt_loc_flat = tgt_loc.view(-1).long()
src_loc_flat = src_loc.view(-1).long()
for index_k in self.buffer:
if index_k.shape[0] != 0:
index_k[tgt_loc_flat] = index_k[src_loc_flat]
def get_buffer(self, layer_id: int) -> torch.Tensor:
if self.pool.layer_transfer_counter is not None:
self.pool.layer_transfer_counter.wait_until(
layer_id - self.pool.start_layer
)
return self.get_broadcastable_buffer(layer_id)
def get_k_and_scale(
self,
layer_id: int,
seq_len_tensor: torch.Tensor,
page_indices: torch.Tensor,
seq_len_sum: int,
max_seq_len: int,
):
buf = self.get_buffer(layer_id)
self.pool.prefetch_kv_buffer(layer_id)
return index_buf_accessor.GetKAndS.execute(
self.pool,
buf,
page_indices=page_indices,
seq_len_tensor=seq_len_tensor,
seq_len_sum=seq_len_sum,
max_seq_len=max_seq_len,
)
def store_quantized(
self,
layer_id: int,
loc: torch.Tensor,
index_k: torch.Tensor,
index_k_scale: torch.Tensor,
) -> None:
self.invalidate(layer_id)
if self.pool._is_layer_owned(layer_id):
super().store_quantized(layer_id, loc, index_k, index_k_scale)
def invalidate(self, layer_id: int) -> None:
if self.remote_layer_id == layer_id:
self.remote_layer_id = None
def get_broadcastable_buffer(self, layer_id: int) -> torch.Tensor:
if self.remote_layer_id != layer_id:
local_idx = layer_id - self.pool.start_layer
src_tensor = (
self.buffer[local_idx] if self.pool._is_layer_owned(layer_id) else None
)
self.pool._broadcast_tensor_from_owner(
self.remote_buffer,
layer_id,
src_tensor=src_tensor,
)
self.remote_layer_id = layer_id
return self.remote_buffer
def state_buf_infos(self):
owned_layer_ids = [
i
for i in range(self.pool.layer_num)
if self.pool._is_layer_owned(self.pool.start_layer + i)
]
data_ptrs = [self.buffer[i].data_ptr() for i in owned_layer_ids]
data_lens = [self.buffer[i].nbytes for i in owned_layer_ids]
item_lens = [self.buffer[i][0].nbytes for i in owned_layer_ids]
return data_ptrs, data_lens, item_lens
def cpu_copy(self, indices):
page_indices = indices[:: self.pool.page_size] // self.pool.page_size
torch.cuda.synchronize()
index_k_cpu = []
chunk_size = self.pool.cpu_offloading_chunk_size
page_chunk_size = max(1, chunk_size // self.pool.page_size)
for layer_id in range(self.pool.layer_num):
index_k_cpu.append([])
if self.buffer[layer_id].shape[0] == 0:
continue
for i in range(0, len(page_indices), page_chunk_size):
chunk_page_indices = page_indices[i : i + page_chunk_size]
idx_cpu = self.buffer[layer_id][chunk_page_indices].to(
"cpu", non_blocking=True
)
index_k_cpu[-1].append(idx_cpu)
torch.cuda.synchronize()
return index_k_cpu
def load_cpu_copy(self, index_k_cpu, indices) -> None:
page_indices = indices[:: self.pool.page_size] // self.pool.page_size
torch.cuda.synchronize()
chunk_size = self.pool.cpu_offloading_chunk_size
page_chunk_size = max(1, chunk_size // self.pool.page_size)
for layer_id in range(self.pool.layer_num):
if self.buffer[layer_id].shape[0] == 0:
continue
for i in range(0, len(page_indices), page_chunk_size):
chunk_page_indices = page_indices[i : i + page_chunk_size]
idx_cpu = index_k_cpu[layer_id][i // page_chunk_size]
assert idx_cpu.shape[0] == len(chunk_page_indices)
idx_chunk = idx_cpu.to(self.buffer[layer_id].device, non_blocking=True)
self.buffer[layer_id][chunk_page_indices] = idx_chunk
torch.cuda.synchronize()
class LayerSplitDSATokenToKVPool(DSATokenToKVPool): class LayerSplitDSATokenToKVPool(DSATokenToKVPool):
"""DSA KV pool that shards layers across CP ranks with owner-broadcast reads.""" """DSA KV pool that shards layers across CP ranks with owner-broadcast reads."""
@@ -203,35 +339,13 @@ class LayerSplitDSATokenToKVPool(DSATokenToKVPool):
self.pending_remote_kv_broadcast = False self.pending_remote_kv_broadcast = False
self._init_layer_broadcast_comm() self._init_layer_broadcast_comm()
def _create_index_buffers(self): def _create_index_key_cache(self) -> IndexKeyCache:
num_pages = (self.index_buf_size + self.page_size + 1) // self.page_size return LayerSplitIndexKeyCache(self, self.index_buf_size)
with (
torch.cuda.use_mem_pool(self.custom_mem_pool)
if self.custom_mem_pool
else nullcontext()
):
self.index_k_with_scale_buffer = [
torch.zeros(
self._index_buffer_shape(
num_pages if self._is_layer_owned(self.start_layer + i) else 0
),
dtype=self.index_k_with_scale_buffer_dtype,
device=self.device,
)
for i in range(self.layer_num)
]
self.remote_index_k_with_scale_buffer = torch.empty(
self._index_buffer_shape(num_pages),
dtype=self.index_k_with_scale_buffer_dtype,
device=self.device,
)
self.remote_index_layer_id: Optional[int] = None
def _clear_buffers(self): def _clear_buffers(self):
del self.kv_buffer del self.kv_buffer
del self.remote_kv_buffer del self.remote_kv_buffer
del self.remote_index_k_with_scale_buffer self.index_key_cache.clear()
del self.index_k_with_scale_buffer
# ---- MLA latent KV: owned-only writes, owner-broadcast reads ---------- # ---- MLA latent KV: owned-only writes, owner-broadcast reads ----------
@@ -418,96 +532,20 @@ class LayerSplitDSATokenToKVPool(DSATokenToKVPool):
if kv_cache.shape[0] == 0: if kv_cache.shape[0] == 0:
continue continue
kv_cache[tgt_loc_flat] = kv_cache[src_loc_flat] kv_cache[tgt_loc_flat] = kv_cache[src_loc_flat]
for index_k in self.index_k_with_scale_buffer: self.index_key_cache.move(tgt_loc, src_loc)
if index_k.shape[0] == 0:
continue
index_k[tgt_loc_flat] = index_k[src_loc_flat]
# ---- DSA indexer buffer: owned-only writes, owner-broadcast reads ----- # ---- DSA indexer buffer: owned-only writes, owner-broadcast reads -----
def get_broadcastable_index_k_with_scale_buffer( def get_broadcastable_index_k_with_scale_buffer(
self, layer_id: int self, layer_id: int
) -> torch.Tensor: ) -> torch.Tensor:
if self.layer_transfer_counter is not None: return self.index_key_cache.get_buffer(layer_id)
self.layer_transfer_counter.wait_until(layer_id - self.start_layer)
return self._get_broadcastable_index_buffer(layer_id)
def get_index_k_continuous(self, layer_id, seq_len, page_indices):
if self.layer_transfer_counter is not None:
self.layer_transfer_counter.wait_until(layer_id - self.start_layer)
buf = self._get_broadcastable_index_buffer(layer_id)
return index_buf_accessor.GetK.execute(
self, buf, seq_len=seq_len, page_indices=page_indices
)
def get_index_k_scale_continuous(self, layer_id, seq_len, page_indices):
if self.layer_transfer_counter is not None:
self.layer_transfer_counter.wait_until(layer_id - self.start_layer)
buf = self._get_broadcastable_index_buffer(layer_id)
return index_buf_accessor.GetS.execute(
self, buf, seq_len=seq_len, page_indices=page_indices
)
def get_index_k_scale_buffer(
self, layer_id, seq_len_tensor, page_indices, seq_len_sum, max_seq_len
):
if self.layer_transfer_counter is not None:
self.layer_transfer_counter.wait_until(layer_id - self.start_layer)
buf = self._get_broadcastable_index_buffer(layer_id)
# Overlap the latent-KV owner-broadcast with the indexer read.
self.prefetch_kv_buffer(layer_id)
return index_buf_accessor.GetKAndS.execute(
self,
buf,
page_indices=page_indices,
seq_len_tensor=seq_len_tensor,
seq_len_sum=seq_len_sum,
max_seq_len=max_seq_len,
)
def set_index_k_scale_buffer(self, layer_id, loc, index_k, index_k_scale) -> None:
self.invalidate_index_buffer_for_layer(layer_id)
if not self._is_layer_owned(layer_id):
return
buf = self.index_k_with_scale_buffer[layer_id - self.start_layer]
index_buf_accessor.SetKAndS.execute(
pool=self, buf=buf, loc=loc, index_k=index_k, index_k_scale=index_k_scale
)
def invalidate_index_buffer_for_layer(self, layer_id: int) -> None: def invalidate_index_buffer_for_layer(self, layer_id: int) -> None:
if self.remote_index_layer_id == layer_id: self.index_key_cache.invalidate(layer_id)
self.remote_index_layer_id = None
def _get_broadcastable_index_buffer(self, layer_id: int) -> torch.Tensor: def _get_broadcastable_index_buffer(self, layer_id: int) -> torch.Tensor:
if self.remote_index_layer_id != layer_id: return self.index_key_cache.get_broadcastable_buffer(layer_id)
local_idx = self._local_layer_idx(layer_id)
src_tensor = (
self.index_k_with_scale_buffer[local_idx]
if self._is_layer_owned(layer_id)
else None
)
self._broadcast_tensor_from_owner(
self.remote_index_k_with_scale_buffer,
layer_id,
src_tensor=src_tensor,
)
self.remote_index_layer_id = layer_id
return self.remote_index_k_with_scale_buffer
def get_state_buf_infos(self):
owned_layer_ids = [
i
for i in range(self.layer_num)
if self._is_layer_owned(self.start_layer + i)
]
data_ptrs = [
self.index_k_with_scale_buffer[i].data_ptr() for i in owned_layer_ids
]
data_lens = [self.index_k_with_scale_buffer[i].nbytes for i in owned_layer_ids]
item_lens = [
self.index_k_with_scale_buffer[i][0].nbytes for i in owned_layer_ids
]
return data_ptrs, data_lens, item_lens
# ---- HiCache CPU offload: skip empty (non-owned) layers --------------- # ---- HiCache CPU offload: skip empty (non-owned) layers ---------------
@@ -529,22 +567,7 @@ class LayerSplitDSATokenToKVPool(DSATokenToKVPool):
kv_cache_cpu[-1].append(kv_cpu) kv_cache_cpu[-1].append(kv_cpu)
current_platform.synchronize() current_platform.synchronize()
page_indices = indices[:: self.page_size] // self.page_size return {"kv": kv_cache_cpu, "index_k": self.index_key_cache.cpu_copy(indices)}
torch.cuda.synchronize()
index_k_cpu = []
page_chunk_size = max(1, chunk_size // self.page_size)
for layer_id in range(self.layer_num):
index_k_cpu.append([])
if self.index_k_with_scale_buffer[layer_id].shape[0] == 0:
continue
for i in range(0, len(page_indices), page_chunk_size):
chunk_page_indices = page_indices[i : i + page_chunk_size]
idx_cpu = self.index_k_with_scale_buffer[layer_id][
chunk_page_indices
].to("cpu", non_blocking=True)
index_k_cpu[-1].append(idx_cpu)
torch.cuda.synchronize()
return {"kv": kv_cache_cpu, "index_k": index_k_cpu}
def load_cpu_copy(self, kv_cache_cpu_dict, indices, mamba_indices=None): def load_cpu_copy(self, kv_cache_cpu_dict, indices, mamba_indices=None):
from sglang.srt.utils import current_platform from sglang.srt.utils import current_platform
@@ -563,19 +586,4 @@ class LayerSplitDSATokenToKVPool(DSATokenToKVPool):
self.kv_buffer[layer_id][chunk_indices] = kv_chunk self.kv_buffer[layer_id][chunk_indices] = kv_chunk
current_platform.synchronize() current_platform.synchronize()
page_indices = indices[:: self.page_size] // self.page_size self.index_key_cache.load_cpu_copy(kv_cache_cpu_dict["index_k"], indices)
index_k_cpu = kv_cache_cpu_dict["index_k"]
torch.cuda.synchronize()
page_chunk_size = max(1, chunk_size // self.page_size)
for layer_id in range(self.layer_num):
if self.index_k_with_scale_buffer[layer_id].shape[0] == 0:
continue
for i in range(0, len(page_indices), page_chunk_size):
chunk_page_indices = page_indices[i : i + page_chunk_size]
idx_cpu = index_k_cpu[layer_id][i // page_chunk_size]
assert idx_cpu.shape[0] == len(chunk_page_indices)
idx_chunk = idx_cpu.to(
self.index_k_with_scale_buffer[layer_id].device, non_blocking=True
)
self.index_k_with_scale_buffer[layer_id][chunk_page_indices] = idx_chunk
torch.cuda.synchronize()
@@ -0,0 +1,149 @@
from __future__ import annotations
from contextlib import nullcontext
from typing import TYPE_CHECKING
import torch
from sglang.kernels.ops.attention.dsa import index_buf_accessor
if TYPE_CHECKING:
from sglang.srt.mem_cache.memory_pool import DSATokenToKVPool
class IndexKeyCache:
def __init__(self, pool: DSATokenToKVPool, index_buf_size: int):
self.pool = pool
num_pages = (index_buf_size + pool.page_size + 1) // pool.page_size
with (
torch.cuda.use_mem_pool(pool.custom_mem_pool)
if pool.custom_mem_pool
else nullcontext()
):
self.buffer = [
torch.zeros(
self._buffer_shape(self._layer_num_pages(i, num_pages)),
dtype=pool.index_k_with_scale_buffer_dtype,
device=pool.device,
)
for i in range(pool.layer_num)
]
def _buffer_shape(self, num_pages: int) -> tuple[int, int]:
pool = self.pool
return (
num_pages,
pool.page_size
* (pool.index_head_dim + pool.index_head_dim // pool.quant_block_size * 4),
)
def _layer_num_pages(self, layer_idx: int, num_pages: int) -> int:
return num_pages
def clear(self) -> None:
del self.buffer
def move(self, tgt_loc: torch.Tensor, src_loc: torch.Tensor) -> None:
if tgt_loc.numel() == 0:
return
tgt_loc_flat = tgt_loc.view(-1).long()
src_loc_flat = src_loc.view(-1).long()
for index_k in self.buffer:
index_k[tgt_loc_flat] = index_k[src_loc_flat]
def get_local_buffer(self, layer_id: int) -> torch.Tensor:
if self.pool.layer_transfer_counter is not None:
self.pool.layer_transfer_counter.wait_until(
layer_id - self.pool.start_layer
)
return self.buffer[layer_id - self.pool.start_layer]
def get_buffer(self, layer_id: int) -> torch.Tensor:
return self.get_local_buffer(layer_id)
def get_k_continuous(self, layer_id: int, seq_len: int, page_indices: torch.Tensor):
buf = self.get_buffer(layer_id)
return index_buf_accessor.GetK.execute(
self.pool, buf, seq_len=seq_len, page_indices=page_indices
)
def get_k_scale_continuous(
self, layer_id: int, seq_len: int, page_indices: torch.Tensor
):
buf = self.get_buffer(layer_id)
return index_buf_accessor.GetS.execute(
self.pool, buf, seq_len=seq_len, page_indices=page_indices
)
def get_k_and_scale(
self,
layer_id: int,
seq_len_tensor: torch.Tensor,
page_indices: torch.Tensor,
seq_len_sum: int,
max_seq_len: int,
):
buf = self.get_buffer(layer_id)
return index_buf_accessor.GetKAndS.execute(
self.pool,
buf,
page_indices=page_indices,
seq_len_tensor=seq_len_tensor,
seq_len_sum=seq_len_sum,
max_seq_len=max_seq_len,
)
def store_quantized(
self,
layer_id: int,
loc: torch.Tensor,
index_k: torch.Tensor,
index_k_scale: torch.Tensor,
) -> None:
buf = self.buffer[layer_id - self.pool.start_layer]
index_buf_accessor.SetKAndS.execute(
pool=self.pool,
buf=buf,
loc=loc,
index_k=index_k,
index_k_scale=index_k_scale,
)
def cpu_copy(self, indices):
# Retracted pages may be reused before resume, so offload index-K with KV.
page_indices = indices[:: self.pool.page_size] // self.pool.page_size
torch.cuda.synchronize()
index_k_cpu = []
chunk_size = self.pool.cpu_offloading_chunk_size
page_chunk_size = max(1, chunk_size // self.pool.page_size)
for layer_id in range(self.pool.layer_num):
index_k_cpu.append([])
for i in range(0, len(page_indices), page_chunk_size):
chunk_page_indices = page_indices[i : i + page_chunk_size]
idx_cpu = self.buffer[layer_id][chunk_page_indices].to(
"cpu", non_blocking=True
)
index_k_cpu[-1].append(idx_cpu)
torch.cuda.synchronize()
return index_k_cpu
def load_cpu_copy(self, index_k_cpu, indices) -> None:
page_indices = indices[:: self.pool.page_size] // self.pool.page_size
torch.cuda.synchronize()
chunk_size = self.pool.cpu_offloading_chunk_size
page_chunk_size = max(1, chunk_size // self.pool.page_size)
for layer_id in range(self.pool.layer_num):
for i in range(0, len(page_indices), page_chunk_size):
chunk_page_indices = page_indices[i : i + page_chunk_size]
idx_cpu = index_k_cpu[layer_id][i // page_chunk_size]
assert idx_cpu.shape[0] == len(chunk_page_indices)
idx_chunk = idx_cpu.to(self.buffer[0].device, non_blocking=True)
self.buffer[layer_id][chunk_page_indices] = idx_chunk
torch.cuda.synchronize()
def state_buf_infos(self):
layer_num = self.pool.layer_num
data_ptrs = [self.buffer[i].data_ptr() for i in range(layer_num)]
data_lens = [self.buffer[i].nbytes for i in range(layer_num)]
item_lens = [self.buffer[i][0].nbytes for i in range(layer_num)]
return data_ptrs, data_lens, item_lens
+20 -125
View File
@@ -38,7 +38,6 @@ import torch
import triton import triton
import triton.language as tl import triton.language as tl
from sglang.kernels.ops.attention.dsa import index_buf_accessor
from sglang.kernels.ops.attention.dsa.quant_k_cache import ( from sglang.kernels.ops.attention.dsa.quant_k_cache import (
quantize_k_cache, quantize_k_cache,
quantize_k_cache_separate, quantize_k_cache_separate,
@@ -59,6 +58,7 @@ from sglang.srt.layers.quantization.fp4_kv_cache_quant_method import (
) )
from sglang.srt.layers.radix_attention import RadixAttention from sglang.srt.layers.radix_attention import RadixAttention
from sglang.srt.mem_cache.allocator.mamba import MambaSlotAllocator from sglang.srt.mem_cache.allocator.mamba import MambaSlotAllocator
from sglang.srt.mem_cache.index_key_cache import IndexKeyCache
from sglang.srt.mem_cache.kv_vmm_backing import KvVmmBufferOwner from sglang.srt.mem_cache.kv_vmm_backing import KvVmmBufferOwner
from sglang.srt.mem_cache.layout.page_major import ( from sglang.srt.mem_cache.layout.page_major import (
build_page_major_mamba_views, build_page_major_mamba_views,
@@ -4378,58 +4378,28 @@ class DSATokenToKVPool(MLATokenToKVPool):
), f"HIP legacy DSA path requires page_size == 1, got {self.page_size}" ), f"HIP legacy DSA path requires page_size == 1, got {self.page_size}"
else: else:
assert self.page_size == 64 assert self.page_size == 64
self._create_index_buffers() self.index_key_cache = self._create_index_key_cache()
self._finalize_allocation_log(size) self._finalize_allocation_log(size)
def _index_buffer_shape(self, num_pages: int) -> tuple[int, int]: def _create_index_key_cache(self) -> IndexKeyCache:
return ( return IndexKeyCache(self, self.index_buf_size)
num_pages,
self.page_size
* (self.index_head_dim + self.index_head_dim // self.quant_block_size * 4),
)
def _create_index_buffers(self): @property
num_pages = (self.index_buf_size + self.page_size + 1) // self.page_size def index_k_with_scale_buffer(self):
with ( # Preserve direct HiCache access while storage lives behind the facade.
torch.cuda.use_mem_pool(self.custom_mem_pool) return self.index_key_cache.buffer
if self.custom_mem_pool
else nullcontext()
):
self.index_k_with_scale_buffer = [
torch.zeros(
# Layout:
# ref: test_attention.py :: kv_cache_cast_to_fp8
# shape: (num_pages, page_size 64 * head_dim 128 + page_size 64 * fp32_nbytes 4)
# data: for page i,
# * buf[i, :page_size * head_dim] for fp8 data
# * buf[i, page_size * head_dim:].view(float32) for scale
self._index_buffer_shape(num_pages),
dtype=self.index_k_with_scale_buffer_dtype,
device=self.device,
)
for _ in range(self.layer_num)
]
def _clear_buffers(self): def _clear_buffers(self):
super()._clear_buffers() super()._clear_buffers()
del self.index_k_with_scale_buffer self.index_key_cache.clear()
def move_kv_cache(self, tgt_loc: torch.Tensor, src_loc: torch.Tensor): def move_kv_cache(self, tgt_loc: torch.Tensor, src_loc: torch.Tensor):
"""Move latent KV and the DSA indexer cache (key + scale) in lockstep.""" """Move latent KV and the DSA indexer cache (key + scale) in lockstep."""
super().move_kv_cache(tgt_loc, src_loc) super().move_kv_cache(tgt_loc, src_loc)
self.index_key_cache.move(tgt_loc, src_loc)
if tgt_loc.numel() == 0:
return
tgt_loc_flat = tgt_loc.view(-1).long()
src_loc_flat = src_loc.view(-1).long()
for index_k in self.index_k_with_scale_buffer:
index_k[tgt_loc_flat] = index_k[src_loc_flat]
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:
if self.layer_transfer_counter is not None: return self.index_key_cache.get_local_buffer(layer_id)
self.layer_transfer_counter.wait_until(layer_id - self.start_layer)
return self.index_k_with_scale_buffer[layer_id - self.start_layer]
def get_index_k_continuous( def get_index_k_continuous(
self, self,
@@ -4437,12 +4407,7 @@ class DSATokenToKVPool(MLATokenToKVPool):
seq_len: int, seq_len: int,
page_indices: torch.Tensor, page_indices: torch.Tensor,
): ):
if self.layer_transfer_counter is not None: return self.index_key_cache.get_k_continuous(layer_id, seq_len, page_indices)
self.layer_transfer_counter.wait_until(layer_id - self.start_layer)
buf = self.index_k_with_scale_buffer[layer_id - self.start_layer]
return index_buf_accessor.GetK.execute(
self, buf, seq_len=seq_len, page_indices=page_indices
)
def get_index_k_scale_continuous( def get_index_k_scale_continuous(
self, self,
@@ -4450,11 +4415,8 @@ class DSATokenToKVPool(MLATokenToKVPool):
seq_len: int, seq_len: int,
page_indices: torch.Tensor, page_indices: torch.Tensor,
): ):
if self.layer_transfer_counter is not None: return self.index_key_cache.get_k_scale_continuous(
self.layer_transfer_counter.wait_until(layer_id - self.start_layer) layer_id, seq_len, page_indices
buf = self.index_k_with_scale_buffer[layer_id - self.start_layer]
return index_buf_accessor.GetS.execute(
self, buf, seq_len=seq_len, page_indices=page_indices
) )
def get_index_k_scale_buffer( def get_index_k_scale_buffer(
@@ -4465,27 +4427,8 @@ class DSATokenToKVPool(MLATokenToKVPool):
seq_len_sum: int, seq_len_sum: int,
max_seq_len: int, max_seq_len: int,
): ):
""" return self.index_key_cache.get_k_and_scale(
Fused method to get both index K and scale data in a single call using Triton. layer_id, seq_len_tensor, page_indices, seq_len_sum, max_seq_len
More efficient than calling get_index_k_continuous and get_index_k_scale_continuous separately.
:param layer_id: Layer index
:param seq_len: Sequence length
:param page_indices: Page indices tensor
:return: tuple of (k_fp8, k_scale) where
k_fp8: (seq_len, index_head_dim), uint8
k_scale: (seq_len, 4), uint8
"""
if self.layer_transfer_counter is not None:
self.layer_transfer_counter.wait_until(layer_id - self.start_layer)
buf = self.index_k_with_scale_buffer[layer_id - self.start_layer]
return index_buf_accessor.GetKAndS.execute(
self,
buf,
page_indices=page_indices,
seq_len_tensor=seq_len_tensor,
seq_len_sum=seq_len_sum,
max_seq_len=max_seq_len,
) )
def set_index_k_scale_buffer( def set_index_k_scale_buffer(
@@ -4495,68 +4438,20 @@ class DSATokenToKVPool(MLATokenToKVPool):
index_k: torch.Tensor, index_k: torch.Tensor,
index_k_scale: torch.Tensor, index_k_scale: torch.Tensor,
) -> None: ) -> None:
buf = self.index_k_with_scale_buffer[layer_id - self.start_layer] self.index_key_cache.store_quantized(layer_id, loc, index_k, index_k_scale)
index_buf_accessor.SetKAndS.execute(
pool=self, buf=buf, loc=loc, index_k=index_k, index_k_scale=index_k_scale
)
def get_cpu_copy(self, indices, mamba_indices=None): def get_cpu_copy(self, indices, mamba_indices=None):
# DSA keeps a page-indexed index_k_with_scale_buffer alongside kv_buffer.
# Retract frees the slots/pages and they get reused by other reqs'
# set_index_k_scale_buffer, so we must offload it here too -- otherwise
# resume restores kv_buffer but leaves foreign index/scale in place and
# DSA attention reads garbage at those token positions.
kv_cache_cpu = super().get_cpu_copy(indices, mamba_indices=mamba_indices) kv_cache_cpu = super().get_cpu_copy(indices, mamba_indices=mamba_indices)
return {"kv": kv_cache_cpu, "index_k": self.index_key_cache.cpu_copy(indices)}
page_indices = indices[:: self.page_size] // self.page_size
torch.cuda.synchronize()
index_k_cpu = []
chunk_size = self.cpu_offloading_chunk_size
page_chunk_size = max(1, chunk_size // self.page_size)
for layer_id in range(self.layer_num):
index_k_cpu.append([])
for i in range(0, len(page_indices), page_chunk_size):
chunk_page_indices = page_indices[i : i + page_chunk_size]
idx_cpu = self.index_k_with_scale_buffer[layer_id][
chunk_page_indices
].to("cpu", non_blocking=True)
index_k_cpu[-1].append(idx_cpu)
torch.cuda.synchronize()
return {"kv": kv_cache_cpu, "index_k": index_k_cpu}
def load_cpu_copy(self, kv_cache_cpu_dict, indices, mamba_indices=None): def load_cpu_copy(self, kv_cache_cpu_dict, indices, mamba_indices=None):
super().load_cpu_copy( super().load_cpu_copy(
kv_cache_cpu_dict["kv"], indices, mamba_indices=mamba_indices kv_cache_cpu_dict["kv"], indices, mamba_indices=mamba_indices
) )
self.index_key_cache.load_cpu_copy(kv_cache_cpu_dict["index_k"], indices)
page_indices = indices[:: self.page_size] // self.page_size
index_k_cpu = kv_cache_cpu_dict["index_k"]
torch.cuda.synchronize()
chunk_size = self.cpu_offloading_chunk_size
page_chunk_size = max(1, chunk_size // self.page_size)
for layer_id in range(self.layer_num):
for i in range(0, len(page_indices), page_chunk_size):
chunk_page_indices = page_indices[i : i + page_chunk_size]
idx_cpu = index_k_cpu[layer_id][i // page_chunk_size]
assert idx_cpu.shape[0] == len(chunk_page_indices)
idx_chunk = idx_cpu.to(
self.index_k_with_scale_buffer[0].device, non_blocking=True
)
self.index_k_with_scale_buffer[layer_id][chunk_page_indices] = idx_chunk
torch.cuda.synchronize()
def get_state_buf_infos(self): def get_state_buf_infos(self):
data_ptrs = [ return self.index_key_cache.state_buf_infos()
self.index_k_with_scale_buffer[i].data_ptr() for i in range(self.layer_num)
]
data_lens = [
self.index_k_with_scale_buffer[i].nbytes for i in range(self.layer_num)
]
item_lens = [
self.index_k_with_scale_buffer[i][0].nbytes for i in range(self.layer_num)
]
return data_ptrs, data_lens, item_lens
def get_kv_size_bytes(self): def get_kv_size_bytes(self):
kv_size_bytes = super().get_kv_size_bytes() kv_size_bytes = super().get_kv_size_bytes()
+7 -26
View File
@@ -35,6 +35,7 @@ from sglang.srt.layers.attention.dsa.utils import (
is_dsa_enable_prefill_cp, is_dsa_enable_prefill_cp,
is_dsa_prefill_cp_round_robin_split, is_dsa_prefill_cp_round_robin_split,
) )
from sglang.srt.layers.attention.index_topk_share import IndexTopKShareState
from sglang.srt.layers.cp.utils import cp_gather_after_forward, is_cp_v2_active from sglang.srt.layers.cp.utils import cp_gather_after_forward, is_cp_v2_active
from sglang.srt.layers.layernorm import RMSNorm from sglang.srt.layers.layernorm import RMSNorm
from sglang.srt.layers.linear import ReplicatedLinear from sglang.srt.layers.linear import ReplicatedLinear
@@ -261,14 +262,7 @@ class DeepseekModelNextN(nn.Module):
hidden_states = cp_split_and_rebuild_data(forward_batch, hidden_states) hidden_states = cp_split_and_rebuild_data(forward_batch, hidden_states)
positions = cp_split_and_rebuild_position(forward_batch, positions) positions = cp_split_and_rebuild_position(forward_batch, positions)
residual = None residual = None
seed_buf = ( index_topk_share = IndexTopKShareState.from_mtp_carry(forward_batch)
forward_batch.spec_info.dsa_seed_topk_capture
if forward_batch.forward_mode.is_extend(include_draft_extend_v2=True)
else None
)
should_update_dsa_topk_indices = (
forward_batch.reuse_dsa_topk_indices or seed_buf is not None
)
with get_global_expert_distribution_recorder().disable_this_region(): with get_global_expert_distribution_recorder().disable_this_region():
hidden_states, residual, topk_indices = self.decoder( hidden_states, residual, topk_indices = self.decoder(
positions, positions,
@@ -276,11 +270,7 @@ class DeepseekModelNextN(nn.Module):
forward_batch, forward_batch,
residual, residual,
zero_allocator, zero_allocator,
prev_topk_indices=( prev_topk_indices=index_topk_share.topk_indices,
forward_batch.spec_info.dsa_topk_indices
if forward_batch.reuse_dsa_topk_indices
else None
),
) )
if not forward_batch.forward_mode.is_idle(): if not forward_batch.forward_mode.is_idle():
if residual is not None: if residual is not None:
@@ -296,7 +286,7 @@ class DeepseekModelNextN(nn.Module):
forward_batch, forward_batch,
torch.cuda.current_stream(), torch.cuda.current_stream(),
) )
if should_update_dsa_topk_indices and topk_indices is not None: if index_topk_share.should_publish and topk_indices is not None:
topk_indices = _gather_dsa_topk_indices_for_cp( topk_indices = _gather_dsa_topk_indices_for_cp(
topk_indices, topk_indices,
local_num_tokens, local_num_tokens,
@@ -306,21 +296,12 @@ class DeepseekModelNextN(nn.Module):
) )
elif ( elif (
cp_v2_active cp_v2_active
and should_update_dsa_topk_indices and index_topk_share.should_publish
and topk_indices is not None and topk_indices is not None
): ):
topk_indices = cp_gather_after_forward(topk_indices, forward_batch) topk_indices = cp_gather_after_forward(topk_indices, forward_batch)
if should_update_dsa_topk_indices and topk_indices is not None: index_topk_share.update(topk_indices)
if forward_batch.reuse_dsa_topk_indices: index_topk_share.publish()
forward_batch.spec_info.dsa_topk_indices = topk_indices
if seed_buf is not None:
sel = forward_batch.spec_info.dsa_seed_topk_select
src = (
topk_indices[: seed_buf.shape[0]]
if sel is None
else topk_indices[sel]
)
seed_buf[: src.shape[0]].copy_(src)
finally: finally:
exit_stack.close() exit_stack.close()
+9 -5
View File
@@ -67,6 +67,7 @@ from sglang.srt.layers.attention.dsa.utils import (
dsa_use_prefill_cp, dsa_use_prefill_cp,
is_dsa_enable_prefill_cp, is_dsa_enable_prefill_cp,
) )
from sglang.srt.layers.attention.index_topk_share import IndexTopKShareState
from sglang.srt.layers.aux_hidden_states import ( from sglang.srt.layers.aux_hidden_states import (
AuxHiddenStateAccumulator, AuxHiddenStateAccumulator,
AuxHiddenStatePacker, AuxHiddenStatePacker,
@@ -2729,18 +2730,21 @@ class DeepseekV2Model(nn.Module):
else: else:
hidden_states = input_embeds hidden_states = input_embeds
residual = None residual = None
initial_topk_indices = None
else: else:
assert pp_proxy_tensors is not None assert pp_proxy_tensors is not None
hidden_states = pp_proxy_tensors["hidden_states"] hidden_states = pp_proxy_tensors["hidden_states"]
residual = pp_proxy_tensors["residual"] residual = pp_proxy_tensors["residual"]
topk_indices = pp_proxy_tensors.tensors.get("topk_indices") initial_topk_indices = pp_proxy_tensors.tensors.get("topk_indices")
index_topk_share = IndexTopKShareState(forward_batch, initial_topk_indices)
if not self.pp_group.is_first_rank:
assert not ( assert not (
not forward_batch.forward_mode.is_idle() not forward_batch.forward_mode.is_idle()
and hidden_states.shape[0] != 0 and hidden_states.shape[0] != 0
and self.use_dsa and self.use_dsa
and dsa_forward_uses_topk and dsa_forward_uses_topk
and dsa_layer_skips_topk(self.config, self.start_layer) and dsa_layer_skips_topk(self.config, self.start_layer)
and topk_indices is None and index_topk_share.topk_indices is None
), ( ), (
f"PP stage starting at layer {self.start_layer} requires DSA " f"PP stage starting at layer {self.start_layer} requires DSA "
"topk_indices from the previous stage." "topk_indices from the previous stage."
@@ -2802,8 +2806,6 @@ class DeepseekV2Model(nn.Module):
normal_end_layer = normal_start_layer = 0 normal_end_layer = normal_start_layer = 0
# Append-compatible, so the shared capture path below is unchanged. # Append-compatible, so the shared capture path below is unchanged.
aux_hidden_states = AuxHiddenStatePacker(len(self.layers_to_capture)) aux_hidden_states = AuxHiddenStatePacker(len(self.layers_to_capture))
if self.pp_group.is_first_rank:
topk_indices = None
for i in range(normal_start_layer, normal_end_layer): for i in range(normal_start_layer, normal_end_layer):
# NOTE: torch dynamo does not support graph break in context manager # NOTE: torch dynamo does not support graph break in context manager
ctx = ( ctx = (
@@ -2821,7 +2823,7 @@ class DeepseekV2Model(nn.Module):
zero_allocator, zero_allocator,
gemm_output_zero_allocator, gemm_output_zero_allocator,
llama_4_scaling, llama_4_scaling,
prev_topk_indices=topk_indices, prev_topk_indices=index_topk_share.topk_indices,
captured_last_layer_outputs=( captured_last_layer_outputs=(
aux_hidden_states if i in self.layers_to_capture else None aux_hidden_states if i in self.layers_to_capture else None
), ),
@@ -2829,6 +2831,7 @@ class DeepseekV2Model(nn.Module):
i i
), ),
) )
index_topk_share.update(topk_indices)
if normal_end_layer != self.end_layer: if normal_end_layer != self.end_layer:
hidden_states, residual = model_forward_maybe_tbo( hidden_states, residual = model_forward_maybe_tbo(
@@ -2855,6 +2858,7 @@ class DeepseekV2Model(nn.Module):
and self.end_layer < self.config.num_hidden_layers and self.end_layer < self.config.num_hidden_layers
and dsa_layer_skips_topk(self.config, self.end_layer) and dsa_layer_skips_topk(self.config, self.end_layer)
): ):
topk_indices = index_topk_share.topk_indices
if ( if (
not forward_batch.forward_mode.is_idle() not forward_batch.forward_mode.is_idle()
and hidden_states.shape[0] != 0 and hidden_states.shape[0] != 0
@@ -18,6 +18,7 @@ from sglang.srt.hardware_backend.npu.graph_runner.eagle_draft_npu_graph_runner i
from sglang.srt.hardware_backend.npu.graph_runner.npu_graph_runner import NPUGraphRunner from sglang.srt.hardware_backend.npu.graph_runner.npu_graph_runner import NPUGraphRunner
from sglang.srt.kv_canary.runner.canary_manager import context_tuple from sglang.srt.kv_canary.runner.canary_manager import context_tuple
from sglang.srt.layers.attention.flashinfer_backend import FlashInferAttnBackend from sglang.srt.layers.attention.flashinfer_backend import FlashInferAttnBackend
from sglang.srt.layers.attention.index_topk_share import IndexTopKShareState
from sglang.srt.layers.attention.tokenspeed_mla_backend import TokenspeedMLABackend from sglang.srt.layers.attention.tokenspeed_mla_backend import TokenspeedMLABackend
from sglang.srt.layers.attention.triton_backend import TritonAttnBackend from sglang.srt.layers.attention.triton_backend import TritonAttnBackend
from sglang.srt.layers.attention.trtllm_mha_backend import TRTLLMHAAttnBackend from sglang.srt.layers.attention.trtllm_mha_backend import TRTLLMHAAttnBackend
@@ -599,14 +600,11 @@ class EagleDraftWorker(EagleDraftWorkerBase):
# Forward multiple steps # Forward multiple steps
scores = None scores = None
if self.index_share_for_mtp_iteration: with IndexTopKShareState.mtp_iteration(
forward_batch.reuse_dsa_topk_indices = True forward_batch,
# Keep the draft-extend seed so step 0 reuses it; else recompute it. enabled=self.index_share_for_mtp_iteration,
if not ( keep_carry_seed=self.seed_dsa_topk_from_draft_extend,
self.seed_dsa_topk_from_draft_extend
and spec_info.dsa_topk_indices is not None
): ):
spec_info.dsa_topk_indices = None
for i in range(self.speculative_num_steps): for i in range(self.speculative_num_steps):
if draft_tokens_topk1 is not None: if draft_tokens_topk1 is not None:
input_ids = topk_index.flatten() input_ids = topk_index.flatten()
@@ -618,11 +616,9 @@ class EagleDraftWorker(EagleDraftWorkerBase):
token_list.append(tree_info[1]) token_list.append(tree_info[1])
parents_list.append(tree_info[2]) parents_list.append(tree_info[2])
# We don't need to run the last forward. we get 1 token from draft prefill and (#spec steps - 1) tokens here
if i == self.speculative_num_steps - 1: if i == self.speculative_num_steps - 1:
break break
# Set inputs
forward_batch.input_ids = input_ids forward_batch.input_ids = input_ids
# Qwen3-MoE MTP uses a fused RoPE + KV-store path whose cache_loc # Qwen3-MoE MTP uses a fused RoPE + KV-store path whose cache_loc
# argument must be contiguous. # argument must be contiguous.
@@ -634,9 +630,6 @@ class EagleDraftWorker(EagleDraftWorkerBase):
forward_batch.out_cache_loc = out_cache_loc[i] forward_batch.out_cache_loc = out_cache_loc[i]
spec_info.hidden_states = hidden_states spec_info.hidden_states = hidden_states
# Run forward under a per-step ForwardContext so the model layer
# reads attn_backends[i] for the i-th draft step, plus a canary
# index context so canary tracks which draft step is active.
canary_index_ctx = ( canary_index_ctx = (
c.with_active_single_forward_manager(i) c.with_active_single_forward_manager(i)
if (c := self.draft_runner.canary_manager) is not None if (c := self.draft_runner.canary_manager) is not None
@@ -650,9 +643,15 @@ class EagleDraftWorker(EagleDraftWorkerBase):
), ),
canary_index_ctx, canary_index_ctx,
): ):
logits_output = self.draft_runner.forward(forward_batch).logits_output logits_output = self.draft_runner.forward(
maybe_detect_nan(logits_output.next_token_logits, f"draft_forward step {i}") forward_batch
maybe_detect_inf(logits_output.next_token_logits, f"draft_forward step {i}") ).logits_output
maybe_detect_nan(
logits_output.next_token_logits, f"draft_forward step {i}"
)
maybe_detect_inf(
logits_output.next_token_logits, f"draft_forward step {i}"
)
if get_spec().speculative_use_rejection_sampling: if get_spec().speculative_use_rejection_sampling:
probs, topk_p, topk_index = sample_draft_proposal( probs, topk_p, topk_index = sample_draft_proposal(
logits_output.next_token_logits, logits_output.next_token_logits,
@@ -662,7 +661,6 @@ class EagleDraftWorker(EagleDraftWorkerBase):
forward_batch.positions.add_(1) forward_batch.positions.add_(1)
elif self.topk == 1 and not _is_hip: elif self.topk == 1 and not _is_hip:
if _is_cuda: if _is_cuda:
# The positions advance is fused into the kernel.
topk_p, topk_index = draft_topk1_postprocess( topk_p, topk_index = draft_topk1_postprocess(
logits_output.next_token_logits, logits_output.next_token_logits,
forward_batch.positions, forward_batch.positions,
@@ -693,10 +691,6 @@ class EagleDraftWorker(EagleDraftWorkerBase):
topk_index = self.hot_token_id[topk_index] topk_index = self.hot_token_id[topk_index]
hidden_states = logits_output.hidden_states hidden_states = logits_output.hidden_states
if self.index_share_for_mtp_iteration:
spec_info.dsa_topk_indices = None
forward_batch.reuse_dsa_topk_indices = False
draft_probs = ( draft_probs = (
torch.stack(draft_probs_list, dim=1) torch.stack(draft_probs_list, dim=1)
if get_spec().speculative_use_rejection_sampling if get_spec().speculative_use_rejection_sampling
@@ -0,0 +1,163 @@
import sys
from types import SimpleNamespace
import pytest
import torch
from sglang.srt.layers.attention.index_topk_share import IndexTopKShareState
from sglang.test.ci.ci_register import register_cpu_ci
register_cpu_ci(est_time=2, suite="base-a-test-cpu")
def _batch(
reuse: bool, carried, *, is_extend: bool = False, seed_buf=None, seed_select=None
) -> SimpleNamespace:
return SimpleNamespace(
reuse_dsa_topk_indices=reuse,
forward_mode=SimpleNamespace(
is_extend=lambda include_draft_extend_v2: is_extend
),
spec_info=SimpleNamespace(
dsa_topk_indices=carried,
dsa_seed_topk_capture=seed_buf,
dsa_seed_topk_select=seed_select,
),
)
def test_mtp_carry_is_empty_when_reuse_disabled():
batch = _batch(reuse=False, carried="old")
state = IndexTopKShareState.from_mtp_carry(batch)
assert state.topk_indices is None
state.update("new")
state.publish()
assert state.topk_indices == "new"
assert batch.spec_info.dsa_topk_indices == "old"
def test_mtp_carry_reads_and_publishes_when_reuse_enabled():
batch = _batch(reuse=True, carried="old")
state = IndexTopKShareState.from_mtp_carry(batch)
assert state.topk_indices == "old"
state.update("new")
state.publish()
assert state.topk_indices == "new"
assert batch.spec_info.dsa_topk_indices == "new"
def test_target_carry_stays_local_without_publish():
batch = _batch(reuse=False, carried="batch")
state = IndexTopKShareState(batch, "layer")
assert state.topk_indices == "layer"
assert batch.spec_info.dsa_topk_indices == "batch"
state.update(None)
assert state.topk_indices is None
assert batch.spec_info.dsa_topk_indices == "batch"
def test_target_none_does_not_fall_back_to_mtp_carry():
batch = _batch(reuse=True, carried="batch")
state = IndexTopKShareState(batch, None)
assert state.topk_indices is None
state.update("layer")
assert state.topk_indices == "layer"
assert batch.spec_info.dsa_topk_indices == "batch"
def test_publish_captures_draft_extend_seed():
seed_buf = torch.zeros(2, 3, dtype=torch.int64)
batch = _batch(reuse=False, carried=None, is_extend=True, seed_buf=seed_buf)
state = IndexTopKShareState.from_mtp_carry(batch)
assert state.should_publish
state.update(torch.arange(12, dtype=torch.int64).view(4, 3))
state.publish()
assert torch.equal(seed_buf, torch.arange(6, dtype=torch.int64).view(2, 3))
assert batch.spec_info.dsa_topk_indices is None
def test_seed_buffer_is_ignored_outside_extend():
seed_buf = torch.zeros(2, 3, dtype=torch.int64)
batch = _batch(reuse=False, carried=None, is_extend=False, seed_buf=seed_buf)
state = IndexTopKShareState.from_mtp_carry(batch)
assert not state.should_publish
state.update(torch.ones(4, 3, dtype=torch.int64))
state.publish()
assert torch.equal(seed_buf, torch.zeros(2, 3, dtype=torch.int64))
def test_mtp_iteration_clears_batch_state():
batch = _batch(reuse=False, carried="stale")
with IndexTopKShareState.mtp_iteration(batch) as state:
assert state is not None
assert batch.reuse_dsa_topk_indices
assert batch.spec_info.dsa_topk_indices is None
batch.spec_info.dsa_topk_indices = "draft-topk"
assert not batch.reuse_dsa_topk_indices
assert batch.spec_info.dsa_topk_indices is None
def test_mtp_iteration_clears_batch_state_on_exception():
batch = _batch(reuse=False, carried=None)
with pytest.raises(RuntimeError, match="draft step blew up"):
with IndexTopKShareState.mtp_iteration(batch):
batch.spec_info.dsa_topk_indices = "draft-topk"
raise RuntimeError("draft step blew up")
assert not batch.reuse_dsa_topk_indices
assert batch.spec_info.dsa_topk_indices is None
def test_disabled_mtp_iteration_is_passthrough():
batch = _batch(reuse=False, carried="untouched")
with IndexTopKShareState.mtp_iteration(batch, enabled=False) as state:
assert state is None
assert not batch.reuse_dsa_topk_indices
assert batch.spec_info.dsa_topk_indices == "untouched"
assert not batch.reuse_dsa_topk_indices
assert batch.spec_info.dsa_topk_indices == "untouched"
def test_mtp_iteration_preserves_draft_extend_seed():
batch = _batch(reuse=False, carried="extend-seed")
with IndexTopKShareState.mtp_iteration(batch, keep_carry_seed=True) as state:
assert state is not None
assert batch.spec_info.dsa_topk_indices == "extend-seed"
assert state.topk_indices == "extend-seed"
assert not batch.reuse_dsa_topk_indices
assert batch.spec_info.dsa_topk_indices is None
def test_mtp_iteration_clears_missing_draft_extend_seed():
batch = _batch(reuse=False, carried=None)
with IndexTopKShareState.mtp_iteration(batch, keep_carry_seed=True):
assert batch.spec_info.dsa_topk_indices is None
assert not batch.reuse_dsa_topk_indices
assert batch.spec_info.dsa_topk_indices is None
if __name__ == "__main__":
sys.exit(pytest.main([__file__]))
@@ -104,8 +104,12 @@ def _run(rank: int, world: int, port: int):
# Indexer buffer owner-broadcast: owner writes a layer-distinct value, then # Indexer buffer owner-broadcast: owner writes a layer-distinct value, then
# every rank must read it back for every layer. # every rank must read it back for every layer.
for layer_id in range(LAYER_NUM): for layer_id in range(LAYER_NUM):
store_buf = pool.get_index_k_with_scale_buffer(layer_id)
assert (
store_buf.data_ptr() == pool.index_k_with_scale_buffer[layer_id].data_ptr()
)
if pool._is_layer_owned(layer_id): if pool._is_layer_owned(layer_id):
pool.index_k_with_scale_buffer[layer_id].fill_(layer_id + 10) store_buf.fill_(layer_id + 10)
torch.cuda.synchronize() torch.cuda.synchronize()
torch.distributed.barrier() torch.distributed.barrier()
for layer_id in range(LAYER_NUM): for layer_id in range(LAYER_NUM):