From c0b6474b43363c2f4bc60fe3d7817d393fb51d32 Mon Sep 17 00:00:00 2001 From: Liangsheng Yin Date: Mon, 17 Aug 2026 16:40:06 -0700 Subject: [PATCH] [Spec] Reduce host-side overhead in ngram draft prep (#35207) --- python/sglang/srt/speculative/ngram_worker.py | 85 ++++++++++++------- 1 file changed, 53 insertions(+), 32 deletions(-) diff --git a/python/sglang/srt/speculative/ngram_worker.py b/python/sglang/srt/speculative/ngram_worker.py index 39caaa791..570533cb9 100644 --- a/python/sglang/srt/speculative/ngram_worker.py +++ b/python/sglang/srt/speculative/ngram_worker.py @@ -241,47 +241,61 @@ class NGRAMWorker(BaseSpecWorker): bs = len(batch.reqs) stride = self.draft_token_num - prev_token_ids, prev_accept_lens = ( - batch.spec_info.accept_tokens, - batch.spec_info.accept_lens, - ) - if not prev_token_ids.is_cpu: - prev_token_ids = prev_token_ids.cpu() - prev_accept_lens = prev_accept_lens.cpu() - # Worker-level staging: written here at draft prep, consumed by - # _update_ngram_corpus after verify within the same forward call. - self.prev_token_ids = prev_token_ids.tolist() - self.prev_accept_lens = prev_accept_lens.tolist() - - self.ngram_corpus.synchronize() - req_ids = [] - batch_tokens = [] - total_lens = [] - assert len(batch.reqs) == len(self.prev_accept_lens) # Overlap mode processes results one iteration behind, so the last # round's accepted tokens are not yet in req.output_ids and must be # spliced in from spec_info. Sync mode and grammar batches process # results before the next draft prep, so output_ids is already # complete and splicing would duplicate the tail. use_prev_tokens = self.enable_overlap and not batch.grammar_needs_sync() - i = 0 - for req in batch.reqs: + + # Accept-independent prep, hoisted above the blocking .cpu() below. + req_ids = [req.rid for req in batch.reqs] + # Only the last max_trie_depth tokens can match; the ids are + # array.array, so list() the tails for list concat below. + input_tails = [ + list(req.origin_input_ids[-self.max_trie_depth :]) for req in batch.reqs + ] + output_tails = [ + list(req.output_ids[-self.max_trie_depth :]) for req in batch.reqs + ] + base_lens = [ + len(req.origin_input_ids) + len(req.output_ids) for req in batch.reqs + ] + + if use_prev_tokens: + prev_token_ids, prev_accept_lens = ( + batch.spec_info.accept_tokens, + batch.spec_info.accept_lens, + ) + if not prev_token_ids.is_cpu: + prev_token_ids = prev_token_ids.cpu() + prev_accept_lens = prev_accept_lens.cpu() + # Worker-level staging: written here at draft prep, consumed by + # _update_ngram_corpus after verify within the same forward call. + self.prev_token_ids = prev_token_ids.tolist() + self.prev_accept_lens = prev_accept_lens.tolist() + assert bs == len(self.prev_accept_lens) + else: + # _update_ngram_corpus still reads the staging; fill it empty. + self.prev_token_ids = [] + self.prev_accept_lens = [0] * bs + + self.ngram_corpus.synchronize() + batch_tokens = [] + total_lens = [] + for i in range(bs): prev_tokens = ( self.prev_token_ids[i * stride : i * stride + self.prev_accept_lens[i]] if use_prev_tokens else [] ) check_token = self._efficient_concat_last_n( - list(req.origin_input_ids), - list(req.output_ids[-self.max_trie_depth :]) + prev_tokens, + input_tails[i], + output_tails[i] + prev_tokens, self.max_trie_depth, ) - req_ids.append(req.rid) batch_tokens.append(check_token) - i += 1 - total_lens.append( - len(req.origin_input_ids) + len(req.output_ids) + len(prev_tokens) - ) + total_lens.append(base_lens[i] + len(prev_tokens)) req_drafts, mask = self.ngram_corpus.batch_get( req_ids, batch_tokens, total_lens ) @@ -309,6 +323,17 @@ class NGRAMWorker(BaseSpecWorker): tree_mask = self.tree_mask_batch[bs] draft_tokens = self.draft_tokens_batch[bs] + # Pre-sync: only needs seq_lens_cpu, so build before the blocking + # accept sync inside _prepare_draft_tokens. + ones_masks = None + if USE_FULL_MASK and not _is_cpu: + ones_masks = [ + torch.ones( + (self.draft_token_num, batch.seq_lens_cpu[i]), device=self.device + ) + for i in range(bs) + ] + req_drafts, mask = self._prepare_draft_tokens(batch) tree_mask.copy_(torch.from_numpy(mask), non_blocking=True) draft_tokens.copy_(torch.from_numpy(req_drafts), non_blocking=True) @@ -335,13 +360,9 @@ class NGRAMWorker(BaseSpecWorker): mask = mask.reshape(bs, self.draft_token_num, self.draft_token_num) # TODO(siyuan): the for loop here leads to significant overhead in large batch size. Can be written into a kernel. for i in range(bs): - seq_len = batch.seq_lens_cpu[i] - req_mask = torch.ones( - (self.draft_token_num, seq_len), device=self.device - ) req_mask = torch.cat( ( - req_mask, + ones_masks[i], torch.from_numpy(mask[i]).to( device=self.device, non_blocking=True ), @@ -393,7 +414,7 @@ class NGRAMWorker(BaseSpecWorker): else [] ) put_ids = self._efficient_concat_last_n( - list(req.origin_input_ids), + list(req.origin_input_ids[-self.max_trie_depth :]), list(req.output_ids[-self.max_trie_depth :]) + prev_tokens, self.max_trie_depth, )