Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
37 changes: 32 additions & 5 deletions swift/infer_engine/transformers_engine.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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)]
Expand All @@ -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]

Expand All @@ -296,38 +312,46 @@ 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]):
if is_finished[i]:
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)
continue
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(
Expand Down Expand Up @@ -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:
Expand Down
11 changes: 10 additions & 1 deletion tests/infer/test_transformers_worker.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand All @@ -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:
Expand Down Expand Up @@ -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__':
Expand Down
Loading