diff --git a/benchmark/mtbench/bench_sglang_eagle.py b/benchmark/mtbench/bench_sglang_eagle.py index 3eb6036c7..7ad3dcb12 100644 --- a/benchmark/mtbench/bench_sglang_eagle.py +++ b/benchmark/mtbench/bench_sglang_eagle.py @@ -38,7 +38,7 @@ def write_answers(filename, model_id, questions, answers): "model_id": model_id, "choices": { "index": 0, - "turns": [answers[i][0], answers[i][1]], + "prompt": [answers[i][0], answers[i][1]], }, "tstamp": time.time(), } @@ -60,7 +60,7 @@ def main(args): # Construct prompts questions = load_questions(args.question_file)[: args.num_questions] arguments = [ - {"question_1": q["turns"][0], "question_2": q["turns"][1]} for q in questions + {"question_1": q["prompt"][0], "question_2": q["prompt"][1]} for q in questions ] # Select backend diff --git a/python/sglang/srt/layers/attention/flashattention_backend.py b/python/sglang/srt/layers/attention/flashattention_backend.py index 9541e3a7f..385c988a1 100644 --- a/python/sglang/srt/layers/attention/flashattention_backend.py +++ b/python/sglang/srt/layers/attention/flashattention_backend.py @@ -14,6 +14,7 @@ from sglang.srt.layers.radix_attention import AttentionType from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode from sglang.srt.server_args import get_global_server_args from sglang.srt.speculative.spec_info import SpecInput +from sglang.srt.utils import get_compiler_backend if TYPE_CHECKING: from sglang.srt.layers.radix_attention import RadixAttention @@ -411,7 +412,6 @@ class FlashAttentionBackend(AttentionBackend): metadata.page_table = forward_batch.req_to_token_pool.req_to_token[ forward_batch.req_pool_indices, : metadata.max_seq_len_k ] - metadata_expand = FlashAttentionMetadata() decode_length = self.speculative_step_id + 1 metadata_expand.cache_seqlens_int32 = torch.full( @@ -645,6 +645,40 @@ class FlashAttentionBackend(AttentionBackend): metadata.page_table[:, self.strided_indices] // self.page_size ) + if ( + self.topk > 1 + and forward_batch.forward_mode.is_decode_or_idle() + and forward_batch.spec_info is not None + ): + # Modifies cache_seqlens_int32 and page_table(B, speculative_num_steps). + last_page_lens = forward_batch.seq_lens % self.page_size + # First attention handles prefix - last_page_len part. + metadata.cache_seqlens_int32 -= last_page_lens # Both (B, ) + + # Second attention handles last_page_len + decode part. + expanded_last_page_lens = last_page_lens.repeat_interleave(self.topk) + self.forward_metadata_spec_decode_expand.cache_seqlens_int32 += ( + expanded_last_page_lens + ) + decode_length = self.speculative_step_id + 1 + expand_page_table = cache_loc[:, :decode_length].clone() + strided_indices_expand = torch.arange( + 0, + decode_length, + self.page_size, + device=self.device, + ) + last_page_lens_broadcast = expanded_last_page_lens.unsqueeze(-1).expand( + -1, expand_page_table.shape[1] + ) + expand_page_table -= last_page_lens_broadcast + expand_page_table = ( + expand_page_table[:, strided_indices_expand] // self.page_size + ) + self.forward_metadata_spec_decode_expand.page_table = ( + expand_page_table.to(torch.int32) + ) + self.forward_metadata = metadata def forward_extend( @@ -798,8 +832,13 @@ class FlashAttentionBackend(AttentionBackend): o, softmax_lse, *rest = result o_expand, softmax_lse_expand, *rest_expand = flash_attn_with_kvcache( q=q.contiguous().view(-1, layer.tp_q_head_num, layer.head_dim), - k_cache=key_cache, - v_cache=value_cache, + # Here metadata_expand.page_table is not divided with page_size. + # This is because we loose the fine control of what token to attend, + # but has to attend to some block completely. + k_cache=key_cache.view(-1, 1, layer.tp_k_head_num, layer.head_dim), + v_cache=value_cache.view( + -1, 1, layer.tp_v_head_num, layer.head_dim + ), page_table=self.forward_metadata_spec_decode_expand.page_table, cache_seqlens=self.forward_metadata_spec_decode_expand.cache_seqlens_int32, cu_seqlens_q=self.forward_metadata_spec_decode_expand.cu_seqlens_q, @@ -1112,7 +1151,6 @@ class FlashAttentionBackend(AttentionBackend): page_table=page_table, cache_seqlens=cache_seqlens, cu_seqlens_q=metadata.cu_seqlens_q, - cu_seqlens_k_new=cu_seqlens_k, max_seqlen_q=max_seqlen_q, softmax_scale=layer.scaling, causal=False if use_cascade_attn else causal, @@ -1344,6 +1382,17 @@ class FlashAttentionBackend(AttentionBackend): ), } + if self.page_size > 1: + # Used for indicing expand page_table + self.draft_decode_metadata_topk_expand["strided_indices_expand"] = ( + torch.arange( + 0, + self.speculative_num_steps, + self.page_size, + device=self.device, + ) + ) + if ( self.speculative_num_draft_tokens is not None and self.speculative_num_draft_tokens > 0 @@ -1778,30 +1827,58 @@ class FlashAttentionBackend(AttentionBackend): # When top k > 1, we need two specific draft decode metadata, and then merge states # 1. The first half of metadata for prefix tokens metadata = self.draft_decode_metadata_topk_normal[bs] + if self.page_size > 1: + # First attention handles seq_lens - last_page_lens if page size > 1. + last_page_lens = seq_lens % self.page_size + seq_lens = seq_lens - last_page_lens + # last_page_lens_cpu = last_page_lens.max().item() + # seq_lens_cpu -= last_page_lens_cpu metadata.cache_seqlens_int32.copy_(seq_lens) # metadata.max_seq_len_q = self.topk, already set in capture - metadata.max_seq_len_k = seq_lens_cpu.max().item() # metadata.cu_seqlens_q already set in capture - metadata.cu_seqlens_k[1:].copy_( - torch.cumsum( - metadata.cache_seqlens_int32, dim=0, dtype=torch.int32 - ) + # metadata.cu_seqlens_k is not needed + + metadata.max_seq_len_k = seq_lens_cpu.max().item() + max_seq_pages = ( + metadata.max_seq_len_k + self.page_size - 1 + ) // self.page_size + strided_indices = self.decode_cuda_graph_metadata["strided_indices"] + strided_indices = strided_indices[:max_seq_pages] + page_table = ( + self.req_to_token[ + req_pool_indices[:, None], # shape [bs, 1] + strided_indices[None, :], # shape [1, max_seq_pages] + ] + // self.page_size ) - - page_table = self.req_to_token[ - req_pool_indices, : metadata.max_seq_len_k - ] - - metadata.page_table[:, : metadata.max_seq_len_k].copy_(page_table) - + metadata.page_table[:, :max_seq_pages].copy_(page_table) # 2. The second half of metadata for draft tokens (per_batch_num_tokens = topk) metadata_expand = self.draft_decode_metadata_topk_expand[bs] decode_length = self.speculative_step_id + 1 # shape: [bs, num_steps, topk] -> [bs x topk, num_steps] cache_loc = out_cache_loc.view(-1, self.speculative_num_steps) - metadata_expand.page_table[: cache_loc.shape[0]].copy_( - cache_loc[:, :decode_length] - ) + if self.page_size > 1: + # Second attention handles last_page_len + decode part. + strided_indices_expand = ( + self.draft_decode_metadata_topk_expand.get( + "strided_indices_expand" + ) + ) + update_draft_decode_set_expand_metadata_with_page_size( + metadata_expand.cache_seqlens_int32, # Modifies + metadata_expand.page_table, # Modifies + cache_loc, + last_page_lens, + strided_indices_expand, + decode_length, + bs, + self.topk, + self.page_size, + ) + else: + metadata_expand.page_table[: cache_loc.shape[0]].copy_( + cache_loc[:, :decode_length] + ) # TODO: Handle local attention metadata for draft decode when llama4 eagle is supported else: # Normal Decode @@ -1860,10 +1937,15 @@ class FlashAttentionBackend(AttentionBackend): metadata.cu_seqlens_k[1:].copy_( torch.cumsum(metadata.cache_seqlens_int32, dim=0, dtype=torch.int32) ) - page_table = self.req_to_token[ - req_pool_indices, : metadata.max_seq_len_k + max_seq_pages = ( + metadata.max_seq_len_k + self.page_size - 1 + ) // self.page_size + page_indices = self.req_to_token[ + req_pool_indices[:, None], + self.decode_cuda_graph_metadata["strided_indices"][:max_seq_pages], ] - metadata.page_table[:, : metadata.max_seq_len_k].copy_(page_table) + page_indices //= self.page_size + metadata.page_table[:, :max_seq_pages].copy_(page_indices) # 2. The second half of metadata for draft tokens (per_batch_num_tokens = topk) metadata_expand = self.target_verify_metadata_topk_expand[bs] @@ -1926,7 +2008,6 @@ class FlashAttentionBackend(AttentionBackend): dtype=torch.int32, ) ) - if self.has_swa: metadata_swa = self.target_verify_metadata_topk_swa[bs] self._init_sliding_window_attn_spec_metadata( @@ -2411,3 +2492,32 @@ def normal_decode_set_metadata( strided_indices[:max_seq_pages][None, :], ] page_table[:, :max_seq_pages].copy_(page_indices // page_size) + + +@torch.compile(dynamic=True, backend=get_compiler_backend()) +def update_draft_decode_set_expand_metadata_with_page_size( + cache_seqlens_int32: torch.Tensor, # Modifies + page_table: torch.Tensor, # Modifies + cache_loc: torch.Tensor, + last_page_lens: torch.Tensor, + strided_indices_expand: torch.Tensor, + decode_length: int, + bs: int, + topk: int, + page_size: int, +): + expanded_last_page_lens = last_page_lens.repeat_interleave(topk) + cache_seqlens_int32.copy_(decode_length + expanded_last_page_lens) + expand_page_table = cache_loc[:, :decode_length].clone() + last_page_lens_broadcast = expanded_last_page_lens.unsqueeze(-1).expand( + -1, expand_page_table.shape[1] + ) + expand_page_table -= last_page_lens_broadcast + expand_page_table = ( + expand_page_table[ + :, strided_indices_expand[: (decode_length + page_size - 1) // page_size] + ] + // page_size + ) + max_seq_pages_expand = (decode_length + page_size - 1) // page_size + page_table[:, :max_seq_pages_expand].copy_(expand_page_table) diff --git a/python/sglang/srt/layers/quantization/kv_cache.py b/python/sglang/srt/layers/quantization/kv_cache.py index 7da2dac17..ef7cbe74e 100644 --- a/python/sglang/srt/layers/quantization/kv_cache.py +++ b/python/sglang/srt/layers/quantization/kv_cache.py @@ -9,7 +9,6 @@ from sglang.srt.layers.quantization.base_config import ( QuantizeMethodBase, ) from sglang.srt.layers.quantization.fp8_kernel import is_fp8_fnuz -from sglang.srt.layers.radix_attention import RadixAttention logger = logging.getLogger(__name__) @@ -45,7 +44,7 @@ class BaseKVCacheMethod(QuantizeMethodBase): def apply(self, layer: torch.nn.Module) -> torch.Tensor: raise RuntimeError(f"{self.__class__.__name__}.apply should not be called.") - def process_weights_after_loading(self, layer: RadixAttention) -> None: + def process_weights_after_loading(self, layer) -> None: if layer.k_scale > 0.0 and layer.v_scale > 0.0: # We prefer to use separate k_scale and v_scale if present k_scale = layer.k_scale.to("cpu").tolist() diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index b99eb4762..9df8ee882 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -1720,7 +1720,7 @@ class ServerArgs: if ( self.speculative_eagle_topk > 1 and self.page_size > 1 - and self.attention_backend != "flashinfer" + and self.attention_backend not in ["flashinfer", "fa3"] ): raise ValueError( "speculative_eagle_topk > 1 with page_size > 1 is unstable and produces incorrect results for paged attention backends. This combination is only supported for the 'flashinfer' backend." diff --git a/python/sglang/srt/speculative/eagle_worker.py b/python/sglang/srt/speculative/eagle_worker.py index 7ab965b20..fd60d358a 100644 --- a/python/sglang/srt/speculative/eagle_worker.py +++ b/python/sglang/srt/speculative/eagle_worker.py @@ -376,6 +376,7 @@ class EAGLEWorker(TpModelWorker): if self.page_size == 1: for req in batch.reqs: req.kv_allocated_len += self.speculative_num_steps * self.topk + # TODO: We only need self.speculative_num_steps - 1 * topk cache loc out_cache_loc, token_to_kv_pool_state_backup = alloc_token_slots( batch.tree_cache, num_seqs * self.speculative_num_steps * self.topk, @@ -403,21 +404,13 @@ class EAGLEWorker(TpModelWorker): # "x" means speculative draft tokens # "." means padded tokens - # TODO(lmzheng): The current implementation is still a fake support - # for page size > 1. In the `assign_draft_cache_locs` below, - # we directly move the indices instead of the real kv cache. - # This only works when the kernel backend runs with page size = 1. - # If the kernel backend runs with page size > 1, we need to - # duplicate the real KV cache. The overhead of duplicating KV - # cache seems okay because the draft KV cache only has one layer. - # see a related copy operation in MHATokenToKVPool::move_kv_cache. - ( prefix_lens, seq_lens, last_loc, self.num_new_pages_per_topk, self.extend_lens, + last_page_lens, ) = get_last_loc_large_page_size_large_top_k( batch.req_to_token_pool.req_to_token, batch.req_pool_indices, @@ -427,9 +420,9 @@ class EAGLEWorker(TpModelWorker): self.page_size, ) prefix_lens_cpu = batch.seq_lens_cpu - last_page_lens = prefix_lens_cpu % self.page_size + last_page_lens_cpu = prefix_lens_cpu % self.page_size num_new_pages_per_topk = ( - last_page_lens + self.speculative_num_steps + self.page_size - 1 + last_page_lens_cpu + self.speculative_num_steps + self.page_size - 1 ) // self.page_size seq_lens_cpu = ( prefix_lens_cpu // self.page_size * self.page_size @@ -450,6 +443,20 @@ class EAGLEWorker(TpModelWorker): ) ) + if self.page_size > 1 and self.topk > 1: + last_page_lens_cumsum = torch.cumsum(last_page_lens, dim=0) + duplicate_cache_len = torch.sum(last_page_lens_cpu).item() * (self.topk - 1) + target_cache_loc = torch.zeros( + duplicate_cache_len, dtype=torch.int32, device=self.device + ) + source_cache_loc = torch.zeros( + duplicate_cache_len, dtype=torch.int32, device=self.device + ) + else: + # When source_cache_loc is not needed, simply skip + duplicate_cache_len = 0 + source_cache_loc, target_cache_loc, last_page_lens_cumsum = None, None, None + assign_draft_cache_locs[(num_seqs,)]( batch.req_pool_indices, batch.req_to_token_pool.req_to_token, @@ -457,16 +464,25 @@ class EAGLEWorker(TpModelWorker): self.extend_lens, self.num_new_pages_per_topk, out_cache_loc, + source_cache_loc, + target_cache_loc, + last_page_lens_cumsum, + duplicate_cache_len, batch.req_to_token_pool.req_to_token.shape[1], self.topk, self.speculative_num_steps, self.page_size, next_power_of_2(num_seqs), - next_power_of_2(self.speculative_num_steps), + next_power_of_2(self.speculative_num_steps + self.page_size), ) if self.page_size > 1 and self.topk > 1: + if duplicate_cache_len > 0: + self.draft_model_runner.token_to_kv_pool.move_kv_cache( + target_cache_loc, source_cache_loc + ) # Remove padded slots + # TODO: We only need self.speculative_num_steps - 1 cache loc out_cache_loc = out_cache_loc[ : num_seqs * self.topk * self.speculative_num_steps ] @@ -581,7 +597,7 @@ class EAGLEWorker(TpModelWorker): ) if self.hot_token_id is not None: topk_index = self.hot_token_id[topk_index] - + # TODO: We only need self.speculative_num_steps - 1 cache loc out_cache_loc = out_cache_loc.reshape( forward_batch.batch_size, self.topk, self.speculative_num_steps ) @@ -1056,4 +1072,11 @@ def get_last_loc_large_page_size_large_top_k( prefix_lens, ) - return prefix_lens, seq_lens, last_loc, num_new_pages_per_topk, extend_lens + return ( + prefix_lens, + seq_lens, + last_loc, + num_new_pages_per_topk, + extend_lens, + last_page_lens, + ) diff --git a/python/sglang/srt/speculative/spec_utils.py b/python/sglang/srt/speculative/spec_utils.py index f20ea377a..5e215156b 100644 --- a/python/sglang/srt/speculative/spec_utils.py +++ b/python/sglang/srt/speculative/spec_utils.py @@ -147,6 +147,10 @@ def assign_draft_cache_locs( extend_lens, num_new_pages_per_topk, out_cache_loc, + source_cache_loc, + target_cache_loc, + last_page_lens_cumsum, + duplicate_cache_len: tl.constexpr, pool_len: tl.constexpr, topk: tl.constexpr, speculative_num_steps: tl.constexpr, @@ -175,44 +179,73 @@ def assign_draft_cache_locs( mask = copy_offset < copy_len data = tl.load(out_cache_ptr + copy_offset, mask=mask) tl.store(token_pool + kv_start + copy_offset, data, mask=mask) - - if page_size == 1 or topk == 1: - return - - # Part 2: Copy the indices for the last partial page - prefix_len = tl.load(seq_lens + pid) - last_page_len = prefix_len % page_size - offsets = tl.arange(0, page_size) - mask = offsets < last_page_len - num_new_pages_per_topk_ = tl.load(num_new_pages_per_topk + pid) - prefix_base = token_pool + prefix_len - last_page_len - - for topk_id in range(topk): - value = tl.load(prefix_base + offsets, mask=mask) - tl.store( - prefix_base + topk_id * num_new_pages_per_topk_ * page_size + offsets, - value, - mask=mask, - ) - - # Part 3: Remove the padding in out_cache_loc - iter_offest = tl.arange(0, iter_upper) - for topk_id in range(topk): - indices = tl.load( - prefix_base - + topk_id * num_new_pages_per_topk_ * page_size - + last_page_len - + iter_offest, - mask=iter_offest < speculative_num_steps, - ) - tl.store( - out_cache_loc - + pid * topk * speculative_num_steps - + topk_id * speculative_num_steps - + iter_offest, - indices, - mask=iter_offest < speculative_num_steps, - ) + if page_size != 1 and topk != 1 and duplicate_cache_len > 0: + # Part 2: Copy indices into source_cache_loc and target_cache_loc + # Expected output: src:[8,9,10,8,9,10...] tgt:[16,17,18,24,25,26...] + prefix_len = tl.load(seq_lens + pid) + last_page_len = prefix_len % page_size + offsets = tl.arange(0, page_size) + mask = offsets < last_page_len + num_new_pages_per_topk_ = tl.load(num_new_pages_per_topk + pid) + prefix_base = token_pool + prefix_len - last_page_len + src_indices = tl.load(prefix_base + offsets, mask=mask) + last_page_lens_cumsum_ = tl.load(last_page_lens_cumsum + pid) + # Skip the first one since no copy is needed + for topk_id in range(1, topk): + tl.store( + source_cache_loc + + (topk - 1) * (last_page_lens_cumsum_ - last_page_len) + + (topk_id - 1) * last_page_len + + offsets, + src_indices, + mask=mask, + ) + tgt_indices = tl.load( + prefix_base + topk_id * num_new_pages_per_topk_ * page_size + offsets, + mask=mask, + ) + tl.store( + target_cache_loc + + (topk - 1) * (last_page_lens_cumsum_ - last_page_len) + + (topk_id - 1) * last_page_len + + offsets, + tgt_indices, + mask=mask, + ) + # Part 3: Copy and remove the used indices for duplication + # speculative_num_steps=5, page_size=4, num_new_pages_per_topk_=2, last_page_len=1 + # - xxxxx .. | - xxxxx .. | + # topk=0 topk=1 + # "-" means prefix tokens + # "x" means speculative draft tokens + # "." means padded tokens + # we only want to copy the "x" part. + iter_offset = tl.arange(0, iter_upper) + for topk_id in range(topk): + mask_upper = iter_offset < (speculative_num_steps + last_page_len) + mask_lower = iter_offset >= last_page_len + combined_mask = mask_upper & mask_lower + indices = tl.load( + prefix_base + + topk_id * num_new_pages_per_topk_ * page_size + + iter_offset, + mask=combined_mask, + other=0, + ) + # Shift from previous batches + ptr_offset = pid * speculative_num_steps * topk + # Subtract last_page_len to fill the gap of duplicated last page tokens. + # For example, token pool is (1, 2, 3, 4 ,5) and last page is 1, + # we write 2, 3, 4 to the front of out_cache_loc. + tl.store( + out_cache_loc + + ptr_offset + + topk_id * speculative_num_steps + - last_page_len + + iter_offset, + indices, + mask=combined_mask, + ) @triton.jit diff --git a/python/sglang/test/attention/test_flashattn_backend.py b/python/sglang/test/attention/test_flashattn_backend.py index 719b4d1b0..a58e19498 100644 --- a/python/sglang/test/attention/test_flashattn_backend.py +++ b/python/sglang/test/attention/test_flashattn_backend.py @@ -3,7 +3,10 @@ import unittest import torch from sglang.srt.configs.model_config import AttentionArch -from sglang.srt.layers.attention.flashattention_backend import FlashAttentionBackend +from sglang.srt.layers.attention.flashattention_backend import ( + FlashAttentionBackend, + update_draft_decode_set_expand_metadata_with_page_size, +) from sglang.srt.layers.attention.torch_native_backend import TorchNativeAttnBackend from sglang.srt.layers.radix_attention import RadixAttention from sglang.srt.mem_cache.memory_pool import MHATokenToKVPool @@ -346,5 +349,71 @@ class TestFlashAttentionBackend(CustomTestCase): self._run_attention_test(ForwardMode.DECODE, q_len=1, page_size=64) +class TestUpdateDraftDecodeSetExpandMetadata(CustomTestCase): + def test_update_draft_decode_set_expand_metadata_with_page_size(self): + bs, topk, decode_length, page_size = 1, 2, 1, 4 + + cases = [ + ( + torch.tensor( + [ + [23, 24], + [31, 32], + ], + dtype=torch.int32, + ), + torch.tensor( + [ + [5], + [7], + ], + dtype=torch.int32, + ), + ), + ( + torch.tensor( + [ + [27, 28], + [35, 36], + ], + dtype=torch.int32, + ), + torch.tensor( + [ + [6], + [8], + ], + dtype=torch.int32, + ), + ), + ] + + last_page_lens = torch.tensor([3], dtype=torch.int32) + strided_indices_expand = torch.arange( + 0, decode_length, page_size, dtype=torch.long + ) + + for cache_loc, expected_page_table in cases: + cache_seqlens_int32 = torch.zeros(bs * topk, dtype=torch.int32) + page_table = torch.zeros(bs * topk, decode_length, dtype=torch.int32) + + update_draft_decode_set_expand_metadata_with_page_size( + cache_seqlens_int32=cache_seqlens_int32, + page_table=page_table, + cache_loc=cache_loc, + last_page_lens=last_page_lens, + strided_indices_expand=strided_indices_expand, + decode_length=decode_length, + bs=bs, + topk=topk, + page_size=page_size, + ) + + expected_cache_seqlens = torch.tensor([4, 4], dtype=torch.int32) + + self.assertTrue(torch.equal(cache_seqlens_int32, expected_cache_seqlens)) + self.assertTrue(torch.equal(page_table, expected_page_table)) + + if __name__ == "__main__": unittest.main() diff --git a/python/sglang/test/speculative/test_spec_utils.py b/python/sglang/test/speculative/test_spec_utils.py new file mode 100644 index 000000000..47085afde --- /dev/null +++ b/python/sglang/test/speculative/test_spec_utils.py @@ -0,0 +1,348 @@ +import unittest + +import numpy as np +import torch + +from sglang.srt.mem_cache.memory_pool import copy_all_layer_kv_cache_tiled +from sglang.srt.speculative.spec_utils import assign_draft_cache_locs +from sglang.srt.utils import next_power_of_2 + +BYTES_PER_TILE = 128 + + +class TestSpecUtils(unittest.TestCase): + + def setUp(self): + self.device = "cuda" if torch.cuda.is_available() else "cpu" + self.data_ptrs = torch.zeros(2, 1, dtype=torch.uint64, device=self.device) + self.k_cache = [ + torch.zeros((100, 1, 1), dtype=torch.float32, device=self.device) + ] + self.v_cache = [ + torch.zeros((100, 1, 1), dtype=torch.float32, device=self.device) + ] + self.k_cache[0][:11, 0, 0] = torch.tensor( + [0.0, 0.1, 0.2, 0.3, 0.4, 0.5, 0.6, 0.7, 0.8, 0.9, 1.0], + dtype=torch.float32, + device=self.device, + ) + self.v_cache[0][:11, 0, 0] = torch.tensor( + [-0.0, -0.1, -0.2, -0.3, -0.4, -0.5, -0.6, -0.7, -0.8, -0.9, -1.0], + dtype=torch.float32, + device=self.device, + ) + self.data_ptrs[0, 0] = self.k_cache[0].data_ptr() + self.data_ptrs[1, 0] = self.v_cache[0].data_ptr() + + self.data_strides = torch.tensor( + [ + np.prod(x.shape[1:]) * x.dtype.itemsize + for x in self.k_cache + self.v_cache + ], + device=self.device, + dtype=torch.int64, + ) + + def test_assign_draft_cache_locs_single_seq(self): + # Testing Setup: req_to_token starting from 4 + # 4,5,6,7,{8,9,10}, 8,9,10 is the last partial page, 3 tokens < page_size=4 + # next kv cache will be stored starting 11,12,13... + device = self.device + num_seqs = 1 + page_size = 4 + speculative_num_steps = 5 + topk = 8 + seq_lens_num = 7 + extend_lens_num = 61 # includes the duplicated last page + req_pool_indices = torch.arange(num_seqs, dtype=torch.int32, device=device) + req_to_token = torch.zeros((num_seqs, 100), dtype=torch.int32, device=device) + req_to_token[0, :seq_lens_num] = torch.tensor( + [4, 5, 6, 7, 8, 9, 10], device=device + ) + seq_lens = torch.tensor([seq_lens_num], dtype=torch.int32, device=device) + extend_lens = torch.tensor([extend_lens_num], dtype=torch.int32, device=device) + num_new_pages_per_topk = torch.tensor([2], dtype=torch.int32, device=device) + out_cache_loc = torch.arange(11, 11 + extend_lens_num, device=device) + last_page_lens = torch.tensor([3], dtype=torch.int32, device=device) + last_page_lens_cumsum = torch.cumsum(last_page_lens, dim=0) + duplicate_cache_len = last_page_lens.sum().item() * (topk - 1) + target_cache_loc = torch.zeros( + duplicate_cache_len, dtype=torch.int32, device=device + ) + source_cache_loc = torch.zeros( + duplicate_cache_len, dtype=torch.int32, device=device + ) + assign_draft_cache_locs[(num_seqs,)]( + req_pool_indices, + req_to_token, + seq_lens, + extend_lens, + num_new_pages_per_topk, + out_cache_loc, + source_cache_loc, + target_cache_loc, + last_page_lens_cumsum, + duplicate_cache_len, + req_to_token.shape[1], + topk, + speculative_num_steps, + page_size, + next_power_of_2(num_seqs), + next_power_of_2(speculative_num_steps + page_size), + ) + + out_cache_loc = out_cache_loc[: num_seqs * topk * speculative_num_steps] + expected_source_cache_loc = torch.tensor( + [8, 9, 10] * (topk - 1), device=device, dtype=torch.int32 + ) + assert torch.allclose(source_cache_loc, expected_source_cache_loc) + + copy_all_layer_kv_cache_tiled[(len(self.data_ptrs),)]( + self.data_ptrs, + self.data_strides, + target_cache_loc, + source_cache_loc, + len(target_cache_loc), + next_power_of_2(len(target_cache_loc)), + BYTES_PER_TILE, + ) + assert torch.allclose( + self.k_cache[0][16:19, 0, 0], + torch.tensor( + [0.8, 0.9, 1.0], + dtype=torch.float32, + device=device, + ), + ) + assert torch.allclose( + self.v_cache[0][16:19, 0, 0], + torch.tensor( + [-0.8, -0.9, -1.0], + dtype=torch.float32, + device=device, + ), + ) + + def test_assign_draft_cache_locs_multi_seq(self): + device = self.device + num_seqs = 3 + page_size = 4 + speculative_num_steps = 5 + topk = 8 + req_pool_indices = torch.arange(num_seqs, dtype=torch.int32, device=device) + req_to_token = torch.zeros((num_seqs, 100), dtype=torch.int32, device=device) + seq_lens = torch.tensor([8, 7, 5], dtype=torch.int32, device=device) + extend_lens = torch.tensor([64, 64, 64], dtype=torch.int32, device=device) + num_new_pages_per_topk = torch.tensor( + [2, 2, 2], dtype=torch.int32, device=device + ) + req_to_token = torch.zeros((num_seqs, 100), dtype=torch.int32, device=device) + req_to_token[0, :8] = torch.tensor([4, 5, 6, 7, 8, 9, 10, 11], device=device) + req_to_token[1, :7] = torch.tensor([4, 5, 6, 7, 8, 9, 10], device=device) + req_to_token[2, :5] = torch.tensor([4, 5, 6, 7, 8], device=device) + last_page_lens = torch.tensor([0, 3, 1], dtype=torch.int32, device=device) + last_page_lens_cumsum = torch.cumsum(last_page_lens, dim=0) + duplicate_cache_len = last_page_lens.sum().item() * (topk - 1) + out_cache_loc = torch.arange( + 12, 12 + torch.sum(extend_lens), dtype=torch.int32, device=device + ) + target_cache_loc = torch.zeros( + duplicate_cache_len, dtype=torch.int32, device=device + ) + source_cache_loc = torch.zeros( + duplicate_cache_len, dtype=torch.int32, device=device + ) + assign_draft_cache_locs[(num_seqs,)]( + req_pool_indices, + req_to_token, + seq_lens, + extend_lens, + num_new_pages_per_topk, + out_cache_loc, + source_cache_loc, + target_cache_loc, + last_page_lens_cumsum, + duplicate_cache_len, + req_to_token.shape[1], + topk, + speculative_num_steps, + page_size, + next_power_of_2(num_seqs), + next_power_of_2(speculative_num_steps + page_size), + ) + out_cache_loc = out_cache_loc[: num_seqs * topk * speculative_num_steps] + # fmt: off + expected_out_cache_loc = torch.tensor([ + 12, 13, 14, 15, 16, + 20, 21, 22, 23, 24, + 28, 29, 30, 31, 32, + 36, 37, 38, 39, 40, + 44, 45, 46, 47, 48, + 52, 53, 54, 55, 56, + 60, 61, 62, 63, 64, + 68, 69, 70, 71, 72, + 76, 77, 78, 79, 80, + 84, 85, 86, 87, 88, + 92, 93, 94, 95, 96, + 100, 101, 102, 103, 104, + 108, 109, 110, 111, 112, + 116, 117, 118, 119, 120, + 124, 125, 126, 127, 128, + 132, 133, 134, 135, 136, + 140, 141, 142, 143, 144, + 148, 149, 150, 151, 152, + 156, 157, 158, 159, 160, + 164, 165, 166, 167, 168, + 172, 173, 174, 175, 176, + 180, 181, 182, 183, 184, + 188, 189, 190, 191, 192, + 196, 197, 198, 199, 200 + ], device=device, dtype=torch.int32) + expected_source_cache_loc = torch.tensor([8, 9, 10] * 7 + [8] * 7, device=device, dtype=torch.int32) + expected_target_cache_loc = torch.tensor([ + 81, 82, 83, 89, 90, 91, 97, 98, 99, 105, 106, 107, 113, 114, + 115, 121, 122, 123, 129, 130, 131, 147, 155, 163, 171, 179, 187, 195 + ], device=device, dtype=torch.int32) + # fmt: on + assert torch.allclose(out_cache_loc, expected_out_cache_loc) + assert torch.allclose(source_cache_loc, expected_source_cache_loc) + assert torch.allclose(target_cache_loc, expected_target_cache_loc) + copy_all_layer_kv_cache_tiled[(len(self.data_ptrs),)]( + self.data_ptrs, + self.data_strides, + target_cache_loc, + source_cache_loc, + len(target_cache_loc), + next_power_of_2(len(target_cache_loc)), + BYTES_PER_TILE, + ) + assert torch.allclose( + self.k_cache[0][81:84, 0, 0], + torch.tensor( + [0.8, 0.9, 1.0], + dtype=torch.float32, + device=device, + ), + ) + assert torch.allclose( + self.v_cache[0][81:84, 0, 0], + torch.tensor( + [-0.8, -0.9, -1.0], + dtype=torch.float32, + device=device, + ), + ) + + def test_assign_draft_cache_locs_page_size_1(self): + # Test to make sure page_size=1 not affected + device = self.device + num_seqs = 1 + page_size = 1 + speculative_num_steps = 5 + topk = 8 + seq_lens_num = 7 + extend_lens_num = topk * speculative_num_steps + req_pool_indices = torch.arange(num_seqs, dtype=torch.int32, device=device) + req_to_token = torch.zeros((num_seqs, 100), dtype=torch.int32, device=device) + req_to_token[0, :seq_lens_num] = torch.tensor( + [4, 5, 6, 7, 8, 9, 10], device=device + ) + seq_lens = torch.tensor([seq_lens_num], dtype=torch.int32, device=device) + extend_lens = torch.tensor([extend_lens_num], dtype=torch.int32, device=device) + num_new_pages_per_topk = torch.tensor([2], dtype=torch.int32, device=device) + out_cache_loc = torch.arange(11, 11 + extend_lens_num, device=device) + last_page_lens = torch.tensor([3], dtype=torch.int32, device=device) + duplicate_cache_len = 0 + target_cache_loc = None + source_cache_loc = None + last_page_lens_cumsum = None + assign_draft_cache_locs[(num_seqs,)]( + req_pool_indices, + req_to_token, + seq_lens, + extend_lens, + num_new_pages_per_topk, + out_cache_loc, + source_cache_loc, + target_cache_loc, + last_page_lens_cumsum, + duplicate_cache_len, + req_to_token.shape[1], + topk, + speculative_num_steps, + page_size, + next_power_of_2(num_seqs), + next_power_of_2(speculative_num_steps + page_size), + ) + out_cache_loc = out_cache_loc[: num_seqs * topk * speculative_num_steps] + expected_out_cache_loc = torch.arange(11, 11 + extend_lens_num, device=device) + assert torch.allclose(out_cache_loc, expected_out_cache_loc) + + def test_assign_draft_cache_locs_page_size_gt_spec_steps(self): + device = self.device + num_seqs = 1 + page_size = 16 + speculative_num_steps = 4 + topk = 3 + seq_lens_num = 12 + pool_len = 256 + req_pool_indices = torch.arange(num_seqs, dtype=torch.int32, device=device) + req_to_token = torch.zeros( + (num_seqs, pool_len), dtype=torch.int32, device=device + ) + req_to_token[0, :seq_lens_num] = torch.arange( + seq_lens_num, dtype=torch.int32, device=device + ) + seq_lens = torch.tensor([seq_lens_num], dtype=torch.int32, device=device) + last_page_len = seq_lens_num % page_size + last_page_lens = torch.tensor([last_page_len], dtype=torch.int32, device=device) + last_page_lens_cumsum = torch.cumsum(last_page_lens, dim=0) + num_new_pages_per_topk_val = ( + last_page_len + speculative_num_steps + page_size - 1 + ) // page_size + num_new_pages_per_topk = torch.tensor( + [num_new_pages_per_topk_val], dtype=torch.int32, device=device + ) + extend_lens_num = num_new_pages_per_topk_val * page_size * topk + extend_lens = torch.tensor([extend_lens_num], dtype=torch.int32, device=device) + out_cache_loc = torch.arange( + 2000, 2000 + extend_lens_num, dtype=torch.int32, device=device + ) + duplicate_cache_len = last_page_lens.sum().item() * (topk - 1) + target_cache_loc = torch.zeros( + duplicate_cache_len, dtype=torch.int32, device=device + ) + source_cache_loc = torch.zeros( + duplicate_cache_len, dtype=torch.int32, device=device + ) + assign_draft_cache_locs[(num_seqs,)]( + req_pool_indices, + req_to_token, + seq_lens, + extend_lens, + num_new_pages_per_topk, + out_cache_loc, + source_cache_loc, + target_cache_loc, + last_page_lens_cumsum, + duplicate_cache_len, + req_to_token.shape[1], + topk, + speculative_num_steps, + page_size, + next_power_of_2(num_seqs), + next_power_of_2(speculative_num_steps + page_size), + ) + trimmed = out_cache_loc[: num_seqs * topk * speculative_num_steps] + expected = [] + for topk_id in range(topk): + start = seq_lens_num + topk_id * num_new_pages_per_topk_val * page_size + expected.append( + req_to_token[0, start : start + speculative_num_steps].clone() + ) + expected_out_cache_loc = torch.cat(expected) + assert torch.allclose(trimmed, expected_out_cache_loc) + + +if __name__ == "__main__": + unittest.main() diff --git a/test/srt/test_eagle_infer_a.py b/test/srt/test_eagle_infer_a.py index eb6813a0d..31efd87e2 100644 --- a/test/srt/test_eagle_infer_a.py +++ b/test/srt/test_eagle_infer_a.py @@ -1,3 +1,4 @@ +import os import unittest import requests @@ -34,6 +35,7 @@ class TestEAGLEEngine(CustomTestCase): "speculative_num_draft_tokens": 8, "mem_fraction_static": 0.7, "cuda_graph_max_bs": 5, + "trust_remote_code": True, } NUM_CONFIGS = 2 @@ -189,21 +191,32 @@ class TestEAGLERadixCache(CustomTestCase): "speculative_draft_model_path": DEFAULT_MODEL_NAME_FOR_TEST_EAGLE3, "speculative_algorithm": "EAGLE3", "speculative_num_steps": 2, - "speculative_eagle_topk": 1, - "speculative_num_draft_tokens": 3, + "speculative_eagle_topk": 2, + "speculative_num_draft_tokens": 5, "mem_fraction_static": 0.7, - "cuda_graph_max_bs": 5, "dtype": "float16", + "trust_remote_code": True, + "attention_backend": "fa3", + "skip_server_warmup": True, + "cuda_graph_max_bs": 5, } def test_correctness(self): + os.environ["SGLANG_ALLOW_OVERWRITE_LONGER_CONTEXT_LEN"] = "1" configs = [ # Basic config self.BASE_CONFIG, - # Chunked prefill - {**self.BASE_CONFIG, "chunked_prefill_size": 64}, # Chunked prefill & Page Size > 1 {**self.BASE_CONFIG, "chunked_prefill_size": 64, "page_size": 4}, + {**self.BASE_CONFIG, "page_size": 4}, + # Preferred by some kernels + {**self.BASE_CONFIG, "page_size": 64}, + # Disable CUDA Graph + { + **self.BASE_CONFIG, + "disable_cuda_graph": True, + "page_size": 4, + }, ] for i, config in enumerate(configs): @@ -212,9 +225,11 @@ class TestEAGLERadixCache(CustomTestCase): engine = sgl.Engine(**config, log_level="info", decode_log_interval=10) try: self._test_acc_length(engine) + self._test_batch_generation(engine) finally: engine.shutdown() print("=" * 100) + del os.environ["SGLANG_ALLOW_OVERWRITE_LONGER_CONTEXT_LEN"] def _test_acc_length(self, engine): warmup_prompt = [ @@ -244,6 +259,29 @@ class TestEAGLERadixCache(CustomTestCase): self.assertGreater(acc_length, 2.5) + def _test_batch_generation(self, engine): + prompts = [ + "Hello, my name is", + "The president of the United States is", + "The capital of France is", + "The future of AI is", + ] + params = {"temperature": 0, "max_new_tokens": 50} + + outputs = engine.generate(prompts, params) + for prompt, output in zip(prompts, outputs): + print(f"Prompt: {prompt}") + print(f"Generated: {output['text']}") + print("-" * 40) + + print(f"{engine.get_server_info()=}") + + avg_spec_accept_length = engine.get_server_info()["internal_states"][0][ + "avg_spec_accept_length" + ] + print(f"{avg_spec_accept_length=}") + self.assertGreater(avg_spec_accept_length, 2.0) + @unittest.skipIf(is_in_ci(), "To reduce the CI execution time.") class TestEAGLEDraftExtend(CustomTestCase):