diff --git a/python/sglang/srt/layers/logprob_processor.py b/python/sglang/srt/layers/logprob_processor.py index 0b7833974..e8bbacfd4 100644 --- a/python/sglang/srt/layers/logprob_processor.py +++ b/python/sglang/srt/layers/logprob_processor.py @@ -115,6 +115,12 @@ def get_token_ids_logprobs_raw( vals.append([]) idxs.append([]) continue + if token_ids is None: + # The sequence's rows still occupy logprobs; step over them. + vals.append([]) + idxs.append([]) + pt += pruned_len + continue token_ids_tensor = torch.tensor(token_ids, dtype=torch.long).to( logprobs.device, non_blocking=True ) @@ -169,10 +175,7 @@ def get_top_logprobs_chunk( Returns: int: Number of remaining tokens to process in next chunk """ - # No sequences in the chunk - if logprobs.shape[0] == 0: - return 0 - + # Empty chunks still walk the slice to emit placeholder entries. max_k = max(logits_metadata.top_logprobs_nums) ret = logprobs.topk(max_k, dim=1) values = ret.values.tolist() @@ -208,13 +211,14 @@ def get_top_logprobs_chunk( idx.append(indices[pt + j][:k]) # Append or extend based on whether the sequence was split across chunks - if len(val) > 0: - if split_pruned_len > 0: - input_top_logprobs_val[-1].extend(val) - input_top_logprobs_idx[-1].extend(idx) - else: - input_top_logprobs_val.append(val) - input_top_logprobs_idx.append(idx) + # Split-sequence continuations extend; everyone else owns a fresh + # (possibly empty) entry. + if split_pruned_len > 0: + input_top_logprobs_val[-1].extend(val) + input_top_logprobs_idx[-1].extend(idx) + else: + input_top_logprobs_val.append(val) + input_top_logprobs_idx.append(idx) pt += pruned_len return next_split_pruned_len @@ -242,11 +246,7 @@ def get_token_ids_logprobs_chunk( Returns: int: Number of remaining tokens to process in next chunk """ - - # No sequences in the chunk - if logprobs.shape[0] == 0: - return 0 - + # Empty chunks still walk the slice to emit placeholder entries. pt = 0 next_split_pruned_len = 0 for n, (token_ids, pruned_len) in enumerate( @@ -280,14 +280,14 @@ def get_token_ids_logprobs_chunk( val.append(logprobs[pt + j, token_ids].tolist()) idx.append(token_ids) - # Append or extend based on whether the sequence was split across chunks - if len(val) > 0: - if split_pruned_len > 0: - input_token_ids_logprobs_val[-1].extend(val) - input_token_ids_logprobs_idx[-1].extend(idx) - else: - input_token_ids_logprobs_val.append(val) - input_token_ids_logprobs_idx.append(idx) + # Split-sequence continuations extend; everyone else owns a fresh + # (possibly empty) entry. + if split_pruned_len > 0: + input_token_ids_logprobs_val[-1].extend(val) + input_token_ids_logprobs_idx[-1].extend(idx) + else: + input_token_ids_logprobs_val.append(val) + input_token_ids_logprobs_idx.append(idx) pt += pruned_len return next_split_pruned_len @@ -541,19 +541,17 @@ class InputLogprobProcessor: chunk_sample_indices = sample_indices[chunk_sample_mask] - start_idx sampled_logits[chunk_sample_mask] = chunk_logits[chunk_sample_indices] - # If there are no input logprobs in this chunk, skip the rest - if chunk_indices.numel() == 0: - continue - + # Zero-logprob-row chunks still need the per-sequence bookkeeping below. # Compute the logprobs of the chunk chunk_input_logprobs = chunk_logits[chunk_indices] chunk_input_logprobs = torch.nn.functional.log_softmax( chunk_input_logprobs, dim=-1 ) - # For each chunk, we need to get the slice of the token_to_seq_idx + # End at the last row inside the chunk; token_to_seq_idx[end_idx] + # belongs to the next chunk and would emit its sequence twice. chunk_slice = slice( - token_to_seq_idx[start_idx], token_to_seq_idx[end_idx] + 1 + token_to_seq_idx[start_idx], token_to_seq_idx[end_idx - 1] + 1 ) # Get the logprob of top-k tokens diff --git a/test/registered/unit/layers/test_logprob_chunk_stitching.py b/test/registered/unit/layers/test_logprob_chunk_stitching.py new file mode 100644 index 000000000..21af55cdc --- /dev/null +++ b/test/registered/unit/layers/test_logprob_chunk_stitching.py @@ -0,0 +1,143 @@ +"""Chunked input-logprob processing must match the non-chunked reference. + +Regression for the cross-chunk stitching accounting: zero-logprob-row +sequences (logprob opt-outs in mixed batches, mid-chunked-prefill segments) +were skipped or double-emitted, drifting the per-request entry counts that +the scheduler asserts on. +""" + +import itertools +import unittest +from types import SimpleNamespace + +import torch + +from sglang.srt.layers.logprob_processor import InputLogprobProcessor +from sglang.test.ci.ci_register import register_cpu_ci +from sglang.test.test_utils import CustomTestCase + +register_cpu_ci(est_time=30, suite="base-a-test-cpu") + +VOCAB = 11 +# Heterogeneous per-sequence parameters; uniform ones hide misalignment. +TOPK_CYCLE = [2, 0, 3] +# [] is a valid probe set distinct from None (opt-out). +TOKEN_IDS_CYCLE = [[0, 3], None, [1], []] + + +def _build_batch(seq_specs, with_token_ids): + """seq_specs: list of (extend_len, logprob_start_len). Mirrors + LogitsProcessor._get_pruned_states for the extend-with-logprobs path.""" + pruned_rows = [] + token_to_seq_idx = [] + sample_indices = [] + input_logprob_indices = [] + pruned_lens = [] + sample_pt = -1 + lp_pt = 0 + for idx, (extend_len, start) in enumerate(seq_specs): + eff_start = start - 1 if extend_len == start else start + rows = extend_len - eff_start + pruned_rows.append(torch.randn(rows, VOCAB, dtype=torch.float32)) + token_to_seq_idx.extend([idx] * rows) + sample_pt += rows + sample_indices.append(sample_pt) + n_lp = extend_len - start + input_logprob_indices.extend([lp_pt + i for i in range(n_lp)]) + lp_pt += rows + pruned_lens.append(n_lp) + token_to_seq_idx.append(len(seq_specs) - 1) + metadata = SimpleNamespace( + extend_return_top_logprob=True, + extend_token_ids_logprob=with_token_ids, + top_logprobs_nums=[TOPK_CYCLE[i % 3] for i in range(len(seq_specs))], + extend_logprob_pruned_lens_cpu=pruned_lens, + extend_input_logprob_token_ids_gpu=torch.zeros( + len(input_logprob_indices), dtype=torch.int64 + ), + token_ids_logprobs=( + [TOKEN_IDS_CYCLE[i % len(TOKEN_IDS_CYCLE)] for i in range(len(seq_specs))] + if with_token_ids + else [None] * len(seq_specs) + ), + ) + return ( + torch.cat(pruned_rows), + torch.tensor(sample_indices, dtype=torch.int64), + torch.tensor(input_logprob_indices, dtype=torch.int64), + token_to_seq_idx, + metadata, + ) + + +def _run(proc, batch, chunked, chunk_size): + pruned_states, sample_indices, input_logprob_indices, t2s, metadata = batch + proc.enable_logprobs_chunk = chunked + proc.logprobs_chunk_size = chunk_size + + def get_logits_fn(states, lm_head, logits_metadata, **kwargs): + return states.float() + + return proc.forward( + pruned_states=pruned_states, + sample_indices=sample_indices, + input_logprob_indices=input_logprob_indices, + token_to_seq_idx=t2s, + lm_head=None, + get_logits_fn=get_logits_fn, + logits_metadata=metadata, + ) + + +class TestLogprobChunkStitching(CustomTestCase): + def _sweep(self, with_token_ids): + torch.manual_seed(0) + proc = InputLogprobProcessor() + # (extend_len, start); start == extend_len is the degenerate + # zero-logprob-row shape. + menu = [(1, 1), (2, 2), (3, 0), (4, 1), (5, 5), (2, 0), (6, 2)] + tried = 0 + for n_seqs in (1, 2, 3, 4): + for combo in itertools.product(menu, repeat=n_seqs): + batch = _build_batch(list(combo), with_token_ids) + # Same unit as the production gate: grid rows, not logprob rows. + total_rows = batch[0].shape[0] + for chunk_size in (1, 2, 3, 5): + if total_rows <= chunk_size: + continue + tried += 1 + ref, ref_sampled = _run(proc, batch, False, 10**9) + got, got_sampled = _run(proc, batch, True, chunk_size) + label = f"specs={list(combo)} chunk={chunk_size}" + self.assertEqual( + ref.input_top_logprobs_val, got.input_top_logprobs_val, label + ) + self.assertEqual( + ref.input_top_logprobs_idx, got.input_top_logprobs_idx, label + ) + if with_token_ids: + self.assertEqual( + ref.input_token_ids_logprobs_val, + got.input_token_ids_logprobs_val, + label, + ) + self.assertEqual( + ref.input_token_ids_logprobs_idx, + got.input_token_ids_logprobs_idx, + label, + ) + torch.testing.assert_close( + ref.input_token_logprobs, got.input_token_logprobs, msg=label + ) + torch.testing.assert_close(ref_sampled, got_sampled, msg=label) + self.assertGreater(tried, 1000) + + def test_top_logprobs_stitching(self): + self._sweep(with_token_ids=False) + + def test_token_ids_logprobs_stitching(self): + self._sweep(with_token_ids=True) + + +if __name__ == "__main__": + unittest.main()