From 52a80ef2583bb6791e34a8d6a1354626ed4e09e9 Mon Sep 17 00:00:00 2001 From: taking-lying-flat <1615405@qq.com> Date: Thu, 17 Sep 2026 08:44:22 +0800 Subject: [PATCH 1/3] fix(infer): complete streamed requests independently within batches --- swift/infer_engine/transformers_engine.py | 43 +++- .../test_transformers_stream_completion.py | 201 ++++++++++++++++++ tests/infer/test_transformers_worker.py | 11 +- 3 files changed, 249 insertions(+), 6 deletions(-) create mode 100644 tests/infer/test_transformers_stream_completion.py diff --git a/swift/infer_engine/transformers_engine.py b/swift/infer_engine/transformers_engine.py index 48acb37069..1c766e4dc6 100644 --- a/swift/infer_engine/transformers_engine.py +++ b/swift/infer_engine/transformers_engine.py @@ -23,6 +23,7 @@ from swift.metrics import Metric from swift.model import get_model_processor from swift.template import Template +from swift.template.utils import StopWordsCriteria from swift.tuners import Swift from swift.utils import get_last_valid_indices, patch_kernels, safe_snapshot_download, to_device from .infer_engine import InferEngine @@ -183,7 +184,9 @@ def _infer_worker(self): finished = True res_list = [None] * len(queue_list) for (queue, loop), res in zip(queue_list, res_list): - asyncio.run_coroutine_threadsafe(queue.put(res), loop) + # A missing chunk is not the end-of-stream sentinel. + if res is not None or finished: + asyncio.run_coroutine_threadsafe(queue.put(res), loop) else: for (queue, loop), res in zip(queue_list, res_list_or_gen): asyncio.run_coroutine_threadsafe(queue.put(res), loop) @@ -268,6 +271,15 @@ def _model_generate(**kwargs): self.template.generate(self.model, **kwargs) generate_kwargs = self.template.prepare_generate_kwargs(generate_kwargs, model=self.model) + # Generation runs ahead in another thread; keep independent stop-word state. + stream_stop_criteria = [ + StopWordsCriteria(criteria.tokenizer, criteria.stop_words, **criteria.tokenizer_kwargs) + for criteria in generate_kwargs.get('stopping_criteria', []) if isinstance(criteria, StopWordsCriteria) + ] + eos_token_ids = generation_config.eos_token_id + if isinstance(eos_token_ids, int): + eos_token_ids = [eos_token_ids] + eos_token_ids = set(eos_token_ids or []) thread = Thread(target=_model_generate, kwargs=generate_kwargs) batch_size = inputs['attention_mask'].shape[0] prompt_token_counts = [self._get_num_tokens(inputs, batch_idx=i) for i in range(batch_size)] @@ -279,10 +291,14 @@ def _model_generate(**kwargs): token_idxs = [0] * batch_size raw_batched_generate_ids = None # or torch.Tensor: [batch_size, seq_len] + num_stream_input_tokens = 0 batched_logprobs = [[] for _ in range(batch_size)] while not all_is_finished: try: batched_tokens = next(streamer) + if raw_batched_generate_ids is None and batched_tokens.ndim == 2: + # HF first streams the prompt (or encoder-decoder start token). + num_stream_input_tokens = batched_tokens.shape[1] if batched_tokens.ndim == 1: batched_tokens = batched_tokens[:, None] @@ -296,6 +312,15 @@ def _model_generate(**kwargs): batched_generate_ids = self.template.get_generate_ids(raw_batched_generate_ids, num_prompt_tokens) self._update_batched_logprobs(batched_logprobs, logits_streamer, batched_generate_ids, request_config.top_logprobs) + new_token_ids = raw_batched_generate_ids[:, num_stream_input_tokens:] + num_generated_tokens = new_token_ids.shape[1] + length_finished = ( + generation_config.max_new_tokens is not None + and num_generated_tokens >= generation_config.max_new_tokens) + stopped = torch.zeros(batched_generate_ids.shape[0], dtype=torch.bool, device=batched_generate_ids.device) + if num_generated_tokens > 0: + for criteria in stream_stop_criteria: + stopped |= criteria(new_token_ids, None) res = [] for i in range(batched_generate_ids.shape[0]): @@ -303,17 +328,20 @@ def _model_generate(**kwargs): res.append(None) continue generate_ids = batched_generate_ids[i] + eos_finished = num_generated_tokens > 0 and new_token_ids[i, -1].item() in eos_token_ids + stop_finished = eos_finished or stopped[i].item() + is_finished[i] = all_is_finished or stop_finished or length_finished # ignore pad_token masks = generate_ids != self.tokenizer.pad_token_id + if eos_finished: + # A terminating EOS still counts as a generated token when EOS == PAD. + masks[-1] = True generate_ids = generate_ids[masks].tolist() logprobs_list = None if batched_logprobs[i]: logprobs_list = [logprobs for m, logprobs in zip(masks, batched_logprobs[i]) if m.item()] - is_finished[i] = ( - all_is_finished or is_finished[i] - or len(generate_ids) > 0 and generate_ids[-1] == self.tokenizer.pad_token_id) delta_text = infer_streamers[i].get_printable_text(generate_ids, is_finished[i]) if not delta_text and not is_finished[i]: res.append(None) @@ -321,13 +349,15 @@ def _model_generate(**kwargs): logprobs = self._get_logprobs(logprobs_list, generate_ids[token_idxs[i]:], request_config.top_logprobs) token_idxs[i] = len(generate_ids) - usage_info = self._get_usage_info(prompt_token_counts[i], len(generate_ids)) + usage_info = self._get_usage_info(prompt_token_counts[i], num_generated_tokens) toolcall = None if is_finished[i]: toolcall = self._get_toolcall( self.template.decode_generate_ids(generate_ids, template_inputs=template_inputs[i])) finish_reason = self._get_finish_reason(generation_config.max_new_tokens, usage_info.completion_tokens, is_finished[i]) + if stop_finished: + finish_reason = 'stop' choices = [ ChatCompletionResponseStreamChoice( @@ -508,6 +538,9 @@ async def _gen_wrapper(): if isinstance(item, Exception): raise item yield item + choices = getattr(item, 'choices', None) + if choices and all(choice.finish_reason is not None for choice in choices): + break return _gen_wrapper() else: diff --git a/tests/infer/test_transformers_stream_completion.py b/tests/infer/test_transformers_stream_completion.py new file mode 100644 index 0000000000..4c07acdc27 --- /dev/null +++ b/tests/infer/test_transformers_stream_completion.py @@ -0,0 +1,201 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +import asyncio +import torch +import unittest +from queue import Queue +from threading import Event +from transformers import GenerationConfig +from transformers.utils import is_torch_npu_available +from types import SimpleNamespace +from unittest.mock import Mock, patch + +from swift.infer_engine import RequestConfig, TransformersEngine +from swift.infer_engine.utils import TokensIteratorStreamer +from swift.template import Template + + +class _Tokenizer: + + def __init__(self, pad_token_id): + self.pad_token_id = pad_token_id + + def decode(self, ids, **kwargs): + if isinstance(ids, int): + ids = [ids] + return ''.join({7: '甲', 8: '乙', 9: '结', 10: '束', 11: 'hel', 12: 'lo'}.get(int(token), '') for token in ids) + + def batch_decode(self, ids, **kwargs): + return [self.decode(row) for row in ids] + + +class _StreamEngine(TransformersEngine): + + def __init__(self, steps, *, pad=0, eos=2, stop_words=(), pause_after=None, logprobs=False): + self.steps = steps + self.pause_after = pause_after + self.paused, self.resume, self.generated, self.stopped = (Event() for _ in range(4)) + self._queue = Queue() + self._task_pool, self._adapters_pool = {}, {} + self._task_thread = None + self.max_batch_size = 0 + self.model_name = 'controlled-stream' + self.model = SimpleNamespace(device=torch.npu.current_device() if is_torch_npu_available() else 'cpu') + self.processor = _Tokenizer(pad) + self._get_toolcall = Mock(return_value=None) + self.config = GenerationConfig( + max_new_tokens=len(steps), eos_token_id=eos, pad_token_id=pad, output_logits=logprobs, num_beams=1) + self.template = SimpleNamespace( + tokenizer=self.tokenizer, + template_meta=SimpleNamespace(stop_words=list(stop_words)), + generate=self._generate, + get_generate_ids=lambda ids, length: ids[:, length:], + decode_generate_ids=self.tokenizer.decode) + self.template.prepare_generate_kwargs = lambda kwargs, **kw: Template.prepare_generate_kwargs( + self.template, kwargs, **kw) + + def _generate(self, model, input_ids, streamer, stopping_criteria, logits_processor=(), **kwargs): + try: + streamer.put(input_ids) + for step, tokens in enumerate(self.steps, 1): + for processor in logits_processor: + processor(input_ids, torch.zeros(len(tokens), 16)) + tokens = torch.tensor(tokens) + input_ids = torch.cat([input_ids, tokens[:, None]], dim=1) + streamer.put(tokens) + # Match HF's ordering: tokens are queued before stopping criteria run. + stopping_criteria(input_ids, None) + if step == self.pause_after: + self.paused.set() + self.resume.wait(timeout=5) + finally: + streamer.end() + self.generated.set() + + def _infer(self, infer_requests, request_config, **kwargs): + inputs = {'input_ids': torch.tensor([[4, 2], [0, 5]]), 'attention_mask': torch.tensor([[1, 1], [0, 1]])} + return self._infer_stream( + inputs, + generation_config=self.config, + adapter_request=None, + request_config=request_config, + template_inputs=[None, None]) + + def _fetch_infer_requests(self): + if self.stopped.is_set(): + raise SystemExit + return super()._fetch_infer_requests() + + def _infer_worker(self): + try: + super()._infer_worker() + except SystemExit: + pass + + +class TestStreamCompletion(unittest.TestCase): + + def test_eos_finishes_once_before_batch_end(self): + for pad, eos, end in ((0, 2, 2), (2, 2, 2), (0, [2, 3], 3)): + with self.subTest(pad=pad, eos=eos): + engine = _StreamEngine([[7, 7], [end, 8], [pad, 8], [pad, 8]], pad=pad, eos=eos, logprobs=True) + chunks = list(engine._infer(['a', 'b'], RequestConfig(stream=True, logprobs=True, top_logprobs=1))) + a_chunks = [batch[0] for batch in chunks if batch[0] is not None] + self.assertEqual([r.choices[0].finish_reason for r in a_chunks], [None, 'stop']) + self.assertEqual(a_chunks[-1].usage.completion_tokens, 2) + self.assertEqual(len(a_chunks[-1].choices[0].logprobs['content']), 1) + self.assertEqual(chunks[1][1].choices[0].finish_reason, None) + self.assertTrue(all(batch[0] is None for batch in chunks[2:])) + self.assertEqual(chunks[-1][1].choices[0].finish_reason, 'length') + self.assertEqual(engine._get_toolcall.call_count, 2) + + def test_pad_alone_is_not_a_stop(self): + engine = _StreamEngine([[7, 7], [0, 8], [8, 8], [2, 8]]) + chunks = list(engine._infer(['a', 'b'], RequestConfig(stream=True))) + self.assertIsNone(chunks[1][0]) + a_chunks = [batch[0] for batch in chunks if batch[0] is not None] + self.assertEqual([r.choices[0].finish_reason for r in a_chunks], [None, None, 'stop']) + self.assertEqual(''.join(r.choices[0].delta.content for r in a_chunks), '甲乙') + + def test_final_event_flushes_buffered_text(self): + engine = _StreamEngine([[11, 7], [12, 8], [2, 8], [0, 8]]) + chunks = list(engine._infer(['a', 'b'], RequestConfig(stream=True))) + self.assertIsNone(chunks[0][0]) + self.assertIsNone(chunks[1][0]) + a_chunks = [batch[0] for batch in chunks if batch[0] is not None] + self.assertEqual(len(a_chunks), 1) + self.assertEqual(a_chunks[0].choices[0].delta.content, 'hello') + self.assertEqual(a_chunks[0].choices[0].finish_reason, 'stop') + + def test_eos_on_first_token_or_at_length_limit(self): + for steps in ([[2, 7], [0, 8]], [[7, 7], [2, 8]]): + with self.subTest(steps=steps): + engine = _StreamEngine(steps) + chunks = list(engine._infer(['a', 'b'], RequestConfig(stream=True))) + finals = [ + batch[0] for batch in chunks + if batch[0] is not None and batch[0].choices[0].finish_reason is not None + ] + self.assertEqual(len(finals), 1) + self.assertEqual(finals[0].choices[0].finish_reason, 'stop') + self.assertEqual(chunks[-1][1].choices[0].finish_reason, 'length') + + def test_sampled_pad_counts_towards_length_limit(self): + engine = _StreamEngine([[7, 7], [0, 8]], eos=None) + chunks = list(engine._infer(['a', 'b'], RequestConfig(stream=True))) + self.assertEqual(chunks[-1][0].choices[0].finish_reason, 'length') + self.assertEqual(chunks[-1][0].usage.completion_tokens, 2) + + def test_streamed_prompt_is_not_a_generated_eos(self): + engine = _StreamEngine([[7, 7], [2, 8]]) + engine.template.get_generate_ids = lambda ids, length: ids + chunks = list(engine._infer(['a', 'b'], RequestConfig(stream=True))) + a_chunks = [batch[0] for batch in chunks if batch[0] is not None] + self.assertEqual([r.choices[0].finish_reason for r in a_chunks], [None, 'stop']) + self.assertEqual(a_chunks[-1].usage.completion_tokens, 2) + + def test_stop_words_have_independent_stream_state(self): + for stop in ('结束', [9, 10]): + with self.subTest(stop=stop): + engine = _StreamEngine([[7, 7], [9, 8], [10, 8], [0, 8]], stop_words=[stop]) + original_next = TokensIteratorStreamer.__next__ + + def read_after_generation(streamer): + self.assertTrue(engine.generated.wait(timeout=2)) + return original_next(streamer) + + # The producer has already reached the stop word before any tokens are consumed. + with patch.object(TokensIteratorStreamer, '__next__', read_after_generation): + chunks = list(engine._infer(['a', 'b'], RequestConfig(stream=True))) + a_chunks = [batch[0] for batch in chunks if batch[0] is not None] + self.assertEqual([r.choices[0].finish_reason for r in a_chunks], [None, None, 'stop']) + self.assertEqual(chunks[2][1].choices[0].finish_reason, None) + self.assertIsNone(chunks[-1][0]) + + +class TestStreamCompletionWorker(unittest.IsolatedAsyncioTestCase): + + async def test_request_closes_while_other_request_is_still_generating(self): + engine = _StreamEngine([[7, 7], [0, 8], [2, 8], [0, 8]], pause_after=3) + config = RequestConfig(stream=True) + try: + a, b = await asyncio.gather(engine.infer_async('a', config), engine.infer_async('b', config)) + self.assertTrue(await asyncio.to_thread(engine.paused.wait, 2)) + a_chunks = await asyncio.wait_for(self.collect(a), timeout=2) + self.assertFalse(engine.generated.is_set()) + self.assertEqual([r.choices[0].finish_reason for r in a_chunks], [None, 'stop']) + self.assertEqual(engine._get_toolcall.call_count, 1) + engine.resume.set() + b_chunks = await asyncio.wait_for(self.collect(b), timeout=2) + self.assertEqual([r.choices[0].finish_reason for r in b_chunks], [None, None, None, 'length']) + self.assertEqual(engine._get_toolcall.call_count, 2) + finally: + engine.resume.set() + await asyncio.to_thread(engine.generated.wait, 2) + engine.stopped.set() + if engine._task_thread is not None: + await asyncio.to_thread(engine._task_thread.join, 2) + self.assertFalse(engine._task_thread.is_alive()) + + @staticmethod + async def collect(stream): + return [chunk async for chunk in stream] diff --git a/tests/infer/test_transformers_worker.py b/tests/infer/test_transformers_worker.py index 4370a4659d..0af62789ee 100644 --- a/tests/infer/test_transformers_worker.py +++ b/tests/infer/test_transformers_worker.py @@ -6,6 +6,7 @@ from transformers import GenerationConfig from swift.infer_engine import AdapterRequest, RequestConfig, TransformersEngine +from swift.utils import shutdown_event_loop_in_daemon class _WorkerEngine(TransformersEngine): @@ -23,6 +24,12 @@ def _fetch_infer_requests(self): raise SystemExit return super()._fetch_infer_requests() + def _infer_worker(self): + try: + super()._infer_worker() + except SystemExit: + pass + def _infer(self, infer_requests, request_config, **kwargs): self.batches.append(infer_requests) if request_config.stream and request_config.num_beams == 2: @@ -175,7 +182,9 @@ def request(): self.assertFalse(engine._task_thread.is_alive()) loop = getattr(engine, '_event_loop', None) if loop is not None: - loop.close() + shutdown_event_loop_in_daemon(engine._event_loop_thread, loop) + self.assertFalse(engine._event_loop_thread.is_alive()) + self.assertTrue(loop.is_closed()) if __name__ == '__main__': From 355bc765f38026b54db30ebe3fd3832842057f3f Mon Sep 17 00:00:00 2001 From: taking-lying-flat <1615405@qq.com> Date: Thu, 17 Sep 2026 08:51:34 +0800 Subject: [PATCH 2/3] refactor(infer): simplify stream completion checks --- swift/infer_engine/transformers_engine.py | 12 +++--------- tests/infer/test_transformers_stream_completion.py | 5 ++++- 2 files changed, 7 insertions(+), 10 deletions(-) diff --git a/swift/infer_engine/transformers_engine.py b/swift/infer_engine/transformers_engine.py index 1c766e4dc6..0cc3d94110 100644 --- a/swift/infer_engine/transformers_engine.py +++ b/swift/infer_engine/transformers_engine.py @@ -314,13 +314,7 @@ def _model_generate(**kwargs): request_config.top_logprobs) new_token_ids = raw_batched_generate_ids[:, num_stream_input_tokens:] num_generated_tokens = new_token_ids.shape[1] - length_finished = ( - generation_config.max_new_tokens is not None - and num_generated_tokens >= generation_config.max_new_tokens) - stopped = torch.zeros(batched_generate_ids.shape[0], dtype=torch.bool, device=batched_generate_ids.device) - if num_generated_tokens > 0: - for criteria in stream_stop_criteria: - stopped |= criteria(new_token_ids, None) + stop_results = [criteria(new_token_ids, None) for criteria in stream_stop_criteria if num_generated_tokens] res = [] for i in range(batched_generate_ids.shape[0]): @@ -329,8 +323,8 @@ def _model_generate(**kwargs): continue generate_ids = batched_generate_ids[i] eos_finished = num_generated_tokens > 0 and new_token_ids[i, -1].item() in eos_token_ids - stop_finished = eos_finished or stopped[i].item() - is_finished[i] = all_is_finished or stop_finished or length_finished + stop_finished = eos_finished or any(stopped[i] for stopped in stop_results) + is_finished[i] = all_is_finished or stop_finished # ignore pad_token masks = generate_ids != self.tokenizer.pad_token_id diff --git a/tests/infer/test_transformers_stream_completion.py b/tests/infer/test_transformers_stream_completion.py index 4c07acdc27..69706853ff 100644 --- a/tests/infer/test_transformers_stream_completion.py +++ b/tests/infer/test_transformers_stream_completion.py @@ -186,7 +186,10 @@ async def test_request_closes_while_other_request_is_still_generating(self): self.assertEqual(engine._get_toolcall.call_count, 1) engine.resume.set() b_chunks = await asyncio.wait_for(self.collect(b), timeout=2) - self.assertEqual([r.choices[0].finish_reason for r in b_chunks], [None, None, None, 'length']) + self.assertTrue(all(r.choices[0].finish_reason is None for r in b_chunks[:-1])) + self.assertEqual(b_chunks[-1].choices[0].finish_reason, 'length') + self.assertEqual(''.join(r.choices[0].delta.content for r in b_chunks), '甲乙乙乙') + self.assertEqual(b_chunks[-1].usage.completion_tokens, 4) self.assertEqual(engine._get_toolcall.call_count, 2) finally: engine.resume.set() From 139d019b08e957a09ed41f2bdb8259a61ee37345 Mon Sep 17 00:00:00 2001 From: taking-lying-flat <1615405@qq.com> Date: Thu, 17 Sep 2026 08:53:26 +0800 Subject: [PATCH 3/3] test(infer): keep changes in existing test files --- .../test_transformers_stream_completion.py | 204 ------------------ 1 file changed, 204 deletions(-) delete mode 100644 tests/infer/test_transformers_stream_completion.py diff --git a/tests/infer/test_transformers_stream_completion.py b/tests/infer/test_transformers_stream_completion.py deleted file mode 100644 index 69706853ff..0000000000 --- a/tests/infer/test_transformers_stream_completion.py +++ /dev/null @@ -1,204 +0,0 @@ -# Copyright (c) ModelScope Contributors. All rights reserved. -import asyncio -import torch -import unittest -from queue import Queue -from threading import Event -from transformers import GenerationConfig -from transformers.utils import is_torch_npu_available -from types import SimpleNamespace -from unittest.mock import Mock, patch - -from swift.infer_engine import RequestConfig, TransformersEngine -from swift.infer_engine.utils import TokensIteratorStreamer -from swift.template import Template - - -class _Tokenizer: - - def __init__(self, pad_token_id): - self.pad_token_id = pad_token_id - - def decode(self, ids, **kwargs): - if isinstance(ids, int): - ids = [ids] - return ''.join({7: '甲', 8: '乙', 9: '结', 10: '束', 11: 'hel', 12: 'lo'}.get(int(token), '') for token in ids) - - def batch_decode(self, ids, **kwargs): - return [self.decode(row) for row in ids] - - -class _StreamEngine(TransformersEngine): - - def __init__(self, steps, *, pad=0, eos=2, stop_words=(), pause_after=None, logprobs=False): - self.steps = steps - self.pause_after = pause_after - self.paused, self.resume, self.generated, self.stopped = (Event() for _ in range(4)) - self._queue = Queue() - self._task_pool, self._adapters_pool = {}, {} - self._task_thread = None - self.max_batch_size = 0 - self.model_name = 'controlled-stream' - self.model = SimpleNamespace(device=torch.npu.current_device() if is_torch_npu_available() else 'cpu') - self.processor = _Tokenizer(pad) - self._get_toolcall = Mock(return_value=None) - self.config = GenerationConfig( - max_new_tokens=len(steps), eos_token_id=eos, pad_token_id=pad, output_logits=logprobs, num_beams=1) - self.template = SimpleNamespace( - tokenizer=self.tokenizer, - template_meta=SimpleNamespace(stop_words=list(stop_words)), - generate=self._generate, - get_generate_ids=lambda ids, length: ids[:, length:], - decode_generate_ids=self.tokenizer.decode) - self.template.prepare_generate_kwargs = lambda kwargs, **kw: Template.prepare_generate_kwargs( - self.template, kwargs, **kw) - - def _generate(self, model, input_ids, streamer, stopping_criteria, logits_processor=(), **kwargs): - try: - streamer.put(input_ids) - for step, tokens in enumerate(self.steps, 1): - for processor in logits_processor: - processor(input_ids, torch.zeros(len(tokens), 16)) - tokens = torch.tensor(tokens) - input_ids = torch.cat([input_ids, tokens[:, None]], dim=1) - streamer.put(tokens) - # Match HF's ordering: tokens are queued before stopping criteria run. - stopping_criteria(input_ids, None) - if step == self.pause_after: - self.paused.set() - self.resume.wait(timeout=5) - finally: - streamer.end() - self.generated.set() - - def _infer(self, infer_requests, request_config, **kwargs): - inputs = {'input_ids': torch.tensor([[4, 2], [0, 5]]), 'attention_mask': torch.tensor([[1, 1], [0, 1]])} - return self._infer_stream( - inputs, - generation_config=self.config, - adapter_request=None, - request_config=request_config, - template_inputs=[None, None]) - - def _fetch_infer_requests(self): - if self.stopped.is_set(): - raise SystemExit - return super()._fetch_infer_requests() - - def _infer_worker(self): - try: - super()._infer_worker() - except SystemExit: - pass - - -class TestStreamCompletion(unittest.TestCase): - - def test_eos_finishes_once_before_batch_end(self): - for pad, eos, end in ((0, 2, 2), (2, 2, 2), (0, [2, 3], 3)): - with self.subTest(pad=pad, eos=eos): - engine = _StreamEngine([[7, 7], [end, 8], [pad, 8], [pad, 8]], pad=pad, eos=eos, logprobs=True) - chunks = list(engine._infer(['a', 'b'], RequestConfig(stream=True, logprobs=True, top_logprobs=1))) - a_chunks = [batch[0] for batch in chunks if batch[0] is not None] - self.assertEqual([r.choices[0].finish_reason for r in a_chunks], [None, 'stop']) - self.assertEqual(a_chunks[-1].usage.completion_tokens, 2) - self.assertEqual(len(a_chunks[-1].choices[0].logprobs['content']), 1) - self.assertEqual(chunks[1][1].choices[0].finish_reason, None) - self.assertTrue(all(batch[0] is None for batch in chunks[2:])) - self.assertEqual(chunks[-1][1].choices[0].finish_reason, 'length') - self.assertEqual(engine._get_toolcall.call_count, 2) - - def test_pad_alone_is_not_a_stop(self): - engine = _StreamEngine([[7, 7], [0, 8], [8, 8], [2, 8]]) - chunks = list(engine._infer(['a', 'b'], RequestConfig(stream=True))) - self.assertIsNone(chunks[1][0]) - a_chunks = [batch[0] for batch in chunks if batch[0] is not None] - self.assertEqual([r.choices[0].finish_reason for r in a_chunks], [None, None, 'stop']) - self.assertEqual(''.join(r.choices[0].delta.content for r in a_chunks), '甲乙') - - def test_final_event_flushes_buffered_text(self): - engine = _StreamEngine([[11, 7], [12, 8], [2, 8], [0, 8]]) - chunks = list(engine._infer(['a', 'b'], RequestConfig(stream=True))) - self.assertIsNone(chunks[0][0]) - self.assertIsNone(chunks[1][0]) - a_chunks = [batch[0] for batch in chunks if batch[0] is not None] - self.assertEqual(len(a_chunks), 1) - self.assertEqual(a_chunks[0].choices[0].delta.content, 'hello') - self.assertEqual(a_chunks[0].choices[0].finish_reason, 'stop') - - def test_eos_on_first_token_or_at_length_limit(self): - for steps in ([[2, 7], [0, 8]], [[7, 7], [2, 8]]): - with self.subTest(steps=steps): - engine = _StreamEngine(steps) - chunks = list(engine._infer(['a', 'b'], RequestConfig(stream=True))) - finals = [ - batch[0] for batch in chunks - if batch[0] is not None and batch[0].choices[0].finish_reason is not None - ] - self.assertEqual(len(finals), 1) - self.assertEqual(finals[0].choices[0].finish_reason, 'stop') - self.assertEqual(chunks[-1][1].choices[0].finish_reason, 'length') - - def test_sampled_pad_counts_towards_length_limit(self): - engine = _StreamEngine([[7, 7], [0, 8]], eos=None) - chunks = list(engine._infer(['a', 'b'], RequestConfig(stream=True))) - self.assertEqual(chunks[-1][0].choices[0].finish_reason, 'length') - self.assertEqual(chunks[-1][0].usage.completion_tokens, 2) - - def test_streamed_prompt_is_not_a_generated_eos(self): - engine = _StreamEngine([[7, 7], [2, 8]]) - engine.template.get_generate_ids = lambda ids, length: ids - chunks = list(engine._infer(['a', 'b'], RequestConfig(stream=True))) - a_chunks = [batch[0] for batch in chunks if batch[0] is not None] - self.assertEqual([r.choices[0].finish_reason for r in a_chunks], [None, 'stop']) - self.assertEqual(a_chunks[-1].usage.completion_tokens, 2) - - def test_stop_words_have_independent_stream_state(self): - for stop in ('结束', [9, 10]): - with self.subTest(stop=stop): - engine = _StreamEngine([[7, 7], [9, 8], [10, 8], [0, 8]], stop_words=[stop]) - original_next = TokensIteratorStreamer.__next__ - - def read_after_generation(streamer): - self.assertTrue(engine.generated.wait(timeout=2)) - return original_next(streamer) - - # The producer has already reached the stop word before any tokens are consumed. - with patch.object(TokensIteratorStreamer, '__next__', read_after_generation): - chunks = list(engine._infer(['a', 'b'], RequestConfig(stream=True))) - a_chunks = [batch[0] for batch in chunks if batch[0] is not None] - self.assertEqual([r.choices[0].finish_reason for r in a_chunks], [None, None, 'stop']) - self.assertEqual(chunks[2][1].choices[0].finish_reason, None) - self.assertIsNone(chunks[-1][0]) - - -class TestStreamCompletionWorker(unittest.IsolatedAsyncioTestCase): - - async def test_request_closes_while_other_request_is_still_generating(self): - engine = _StreamEngine([[7, 7], [0, 8], [2, 8], [0, 8]], pause_after=3) - config = RequestConfig(stream=True) - try: - a, b = await asyncio.gather(engine.infer_async('a', config), engine.infer_async('b', config)) - self.assertTrue(await asyncio.to_thread(engine.paused.wait, 2)) - a_chunks = await asyncio.wait_for(self.collect(a), timeout=2) - self.assertFalse(engine.generated.is_set()) - self.assertEqual([r.choices[0].finish_reason for r in a_chunks], [None, 'stop']) - self.assertEqual(engine._get_toolcall.call_count, 1) - engine.resume.set() - b_chunks = await asyncio.wait_for(self.collect(b), timeout=2) - self.assertTrue(all(r.choices[0].finish_reason is None for r in b_chunks[:-1])) - self.assertEqual(b_chunks[-1].choices[0].finish_reason, 'length') - self.assertEqual(''.join(r.choices[0].delta.content for r in b_chunks), '甲乙乙乙') - self.assertEqual(b_chunks[-1].usage.completion_tokens, 4) - self.assertEqual(engine._get_toolcall.call_count, 2) - finally: - engine.resume.set() - await asyncio.to_thread(engine.generated.wait, 2) - engine.stopped.set() - if engine._task_thread is not None: - await asyncio.to_thread(engine._task_thread.join, 2) - self.assertFalse(engine._task_thread.is_alive()) - - @staticmethod - async def collect(stream): - return [chunk async for chunk in stream]