diff --git a/swift/infer_engine/transformers_engine.py b/swift/infer_engine/transformers_engine.py index 48acb37069..0cc3d94110 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,9 @@ 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] + 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]): @@ -303,17 +322,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 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 + 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 +343,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 +532,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_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__':