[Perf][Spec Decoding] Skip cat/topk/sort/gather in draft_forward for topk=1 (#26424)
This commit is contained in:
@@ -0,0 +1,94 @@
|
||||
"""Equivalence tests for the EagleDraftWorker topk=1 chain fast path.
|
||||
|
||||
For topk=1 the draft tree degenerates to a chain, so `draft_forward` skips the
|
||||
cat/topk/sort/gather of the slow path and returns pre-allocated constants. These
|
||||
tests check that the pre-allocated `parent_list` / `top_scores_index` match the
|
||||
slow path (`organize_draft_results`) for num_steps in {1, 2, 3, 4}.
|
||||
"""
|
||||
|
||||
import unittest
|
||||
from types import SimpleNamespace
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.srt.speculative.eagle_utils import organize_draft_results
|
||||
from sglang.srt.speculative.eagle_worker_v2 import EagleDraftWorker
|
||||
from sglang.srt.utils import get_device
|
||||
from sglang.test.ci.ci_register import register_cuda_ci
|
||||
from sglang.test.test_utils import CustomTestCase
|
||||
|
||||
register_cuda_ci(est_time=20, suite="base-b-test-1-gpu-small")
|
||||
|
||||
DEVICE = get_device()
|
||||
|
||||
|
||||
def _make_chain_lists(num_steps: int, bs: int):
|
||||
"""Build the (score, token, parents) lists a topk=1 chain produces.
|
||||
|
||||
Shapes/values mirror `select_top_k_tokens` for topk=1: each step yields one
|
||||
token; the first step's parents are [-1, 0], later steps' parents are [i].
|
||||
"""
|
||||
score_list, token_list, parents_list = [], [], []
|
||||
for i in range(num_steps):
|
||||
# Strictly decreasing scores, as a real chain produces (cumulative probs).
|
||||
score_list.append(torch.full((bs, 1, 1), float(num_steps - i), device=DEVICE))
|
||||
token_list.append(
|
||||
torch.arange(i * bs, (i + 1) * bs, device=DEVICE).unsqueeze(1)
|
||||
)
|
||||
if i == 0:
|
||||
parents_list.append(
|
||||
torch.tensor([-1, 0], dtype=torch.long, device=DEVICE).repeat(bs, 1)
|
||||
)
|
||||
else:
|
||||
parents_list.append(torch.full((bs, 1), i, dtype=torch.long, device=DEVICE))
|
||||
return score_list, token_list, parents_list
|
||||
|
||||
|
||||
def _make_worker(num_steps: int, num_draft_tokens: int):
|
||||
worker = object.__new__(EagleDraftWorker)
|
||||
worker.topk = 1
|
||||
worker.device = DEVICE
|
||||
worker.speculative_num_steps = num_steps
|
||||
worker.speculative_num_draft_tokens = num_draft_tokens
|
||||
worker.server_args = SimpleNamespace(cuda_graph_max_bs=8, max_running_requests=8)
|
||||
return worker
|
||||
|
||||
|
||||
class TestEagleWorkerV2Topk1FastPath(CustomTestCase):
|
||||
def test_fast_path_matches_slow_path(self):
|
||||
bs = 3
|
||||
for num_steps in (1, 2, 3, 4):
|
||||
with self.subTest(num_steps=num_steps):
|
||||
num_draft_tokens = num_steps + 1
|
||||
worker = _make_worker(num_steps, num_draft_tokens)
|
||||
worker._rebuild_topk1_chain_buffers()
|
||||
|
||||
score_list, token_list, parents_list = _make_chain_lists(num_steps, bs)
|
||||
ref_parent, ref_index, ref_tokens = organize_draft_results(
|
||||
score_list, token_list, parents_list, num_draft_tokens
|
||||
)
|
||||
|
||||
fast_parent = worker._topk1_parents_prealloc[:bs]
|
||||
fast_index = worker._topk1_score_indices_prealloc[:bs]
|
||||
fast_tokens = torch.cat(token_list, dim=1)
|
||||
|
||||
self.assertEqual(fast_parent.shape, ref_parent.shape)
|
||||
self.assertEqual(fast_parent.tolist(), ref_parent.long().tolist())
|
||||
self.assertEqual(fast_index.tolist(), ref_index.long().tolist())
|
||||
self.assertEqual(fast_tokens.tolist(), ref_tokens.tolist())
|
||||
|
||||
# The kernel reads these via data_ptr() as contiguous int64.
|
||||
self.assertEqual(fast_parent.dtype, torch.long)
|
||||
self.assertEqual(fast_index.dtype, torch.long)
|
||||
self.assertTrue(fast_parent.is_contiguous())
|
||||
self.assertTrue(fast_index.is_contiguous())
|
||||
|
||||
def test_assert_on_inconsistent_steps_and_draft_tokens(self):
|
||||
# num_draft_tokens must equal num_steps + 1 for topk=1.
|
||||
worker = _make_worker(num_steps=3, num_draft_tokens=3)
|
||||
with self.assertRaises(AssertionError):
|
||||
worker._rebuild_topk1_chain_buffers()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
Reference in New Issue
Block a user