[Logprob] Serve input-logprob temporaries from CUDA-graph-pool dead space (#40038)
Co-authored-by: cctry <csycfl@gmail.com>
This commit is contained in:
@@ -8,9 +8,11 @@ the scheduler asserts on.
|
||||
|
||||
import unittest
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import Mock
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.srt.layers.logits_processor import LogitsProcessorOutput
|
||||
from sglang.srt.layers.logprob_processor import InputLogprobProcessor
|
||||
from sglang.test.ci.ci_register import register_cpu_ci
|
||||
from sglang.test.logprob_test_utils import coverage_cases
|
||||
@@ -52,6 +54,8 @@ def _build_batch(seq_specs, with_token_ids):
|
||||
lp_pt += rows
|
||||
pruned_lens.append(n_lp)
|
||||
metadata = SimpleNamespace(
|
||||
sample_indices_cpu=sample_indices,
|
||||
input_logprob_indices_cpu=input_logprob_indices,
|
||||
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))],
|
||||
@@ -136,6 +140,24 @@ class TestLogprobChunkStitching(CustomTestCase):
|
||||
def test_token_ids_logprobs_stitching(self):
|
||||
self._sweep(with_token_ids=True)
|
||||
|
||||
def test_finalizing_input_logprobs_preserves_request_boundaries(self):
|
||||
rows = [torch.tensor([[1.0], [2.0]]), torch.tensor([[3.0]])]
|
||||
copy_done = Mock()
|
||||
output = LogitsProcessorOutput(
|
||||
next_token_logits=None,
|
||||
input_token_ids_logprobs_val=[[rows[0][:1], rows[0][1:]], [rows[1]]],
|
||||
input_logprobs_copy_done=copy_done,
|
||||
)
|
||||
output.finalize_input_logprobs()
|
||||
self.assertEqual(output.input_token_ids_logprobs_val, [[[1.0], [2.0]], [[3.0]]])
|
||||
copy_done.synchronize.assert_called_once_with()
|
||||
self.assertIsNone(output.input_logprobs_copy_done)
|
||||
|
||||
# Multi-item scoring returns one tensor per request, with no borrow event.
|
||||
output.input_token_ids_logprobs_val = rows
|
||||
output.finalize_input_logprobs()
|
||||
self.assertIs(output.input_token_ids_logprobs_val, rows)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
|
||||
@@ -56,6 +56,8 @@ def _build_batch(seq_specs, dtype, vocab=VOCAB):
|
||||
lp_pt += rows
|
||||
pruned_lens.append(n_lp)
|
||||
metadata = SimpleNamespace(
|
||||
sample_indices_cpu=sample_indices,
|
||||
input_logprob_indices_cpu=input_logprob_indices,
|
||||
extend_return_top_logprob=True,
|
||||
extend_token_ids_logprob=True,
|
||||
top_logprobs_nums=[TOPK_CYCLE[i % 3] for i in range(len(seq_specs))],
|
||||
|
||||
Reference in New Issue
Block a user