diff --git a/batchgen/batchgen_worker.py b/batchgen/batchgen_worker.py index cef5bd2f..165279ba 100644 --- a/batchgen/batchgen_worker.py +++ b/batchgen/batchgen_worker.py @@ -1682,9 +1682,10 @@ def _report_completion(self, uuid: str, gathered_text: str = None) -> None: # Use gathered text if provided, otherwise read from local buffer text = gathered_text if gathered_text is not None else "" if text == "" and seq.decoded_tokens is not None and seq.decoded_length > 0: - token_ids = seq.decoded_tokens[0, :seq.decoded_length].tolist() try: - text = self.tokenizer.decode(token_ids) + text = self._decode_tokens_to_string( + seq.decoded_tokens[:, :seq.decoded_length] + ) except Exception: text = "" self._response_queue.put({ @@ -1693,7 +1694,10 @@ def _report_completion(self, uuid: str, gathered_text: str = None) -> None: "batch_id": getattr(seq, "batch_id", None), "global_idx": seq.global_idx, "text": text, - "prompt_length": seq.prompt_length, + # Re-entry turns prior completion tokens into the internal prompt so + # the KV state can be recomputed. OpenAI usage must still report the + # caller's original prompt, not that internal reconstructed length. + "prompt_length": seq.original_prompt_length, "decoded_length": seq.decoded_length, "finish_reason": self._get_finish_reason(seq), }) @@ -1718,9 +1722,11 @@ def _gather_completed_tokens(self, completed_uuids: List[str]) -> dict: local_idx = self._uuid_to_local_map[uuid] seq = self.global_batch.get_sequence(uuid) if seq is not None and local_idx in self.query_book: - token_ids = self.query_book[local_idx].decoded_tokens[0, :seq.decoded_length].tolist() + decoded_tokens = self.query_book[local_idx].decoded_tokens[ + :, :seq.decoded_length + ] try: - text = self.tokenizer.decode(token_ids) + text = self._decode_tokens_to_string(decoded_tokens) except Exception: text = "" my_tokens[uuid] = text diff --git a/tests/test_completion_usage_accounting.py b/tests/test_completion_usage_accounting.py new file mode 100644 index 00000000..6c3c5a15 --- /dev/null +++ b/tests/test_completion_usage_accounting.py @@ -0,0 +1,109 @@ +from queue import SimpleQueue +from types import SimpleNamespace + +import torch + +import batchgen.batchgen_worker as worker_module +from batchgen.batchgen_worker import BatchGenWorker +from batchgen.sequence import SequenceEntry, SequenceStatus + + +def test_report_completion_keeps_original_prompt_length_after_reentry(): + seq = SequenceEntry( + "reentered", + global_idx=7, + prompt_length=100, + max_decode_length=512, + ) + seq.prompt_length = 228 + seq.decoded_length = 128 + seq.status = SequenceStatus.COMPLETED + + response_queue = SimpleQueue() + worker = SimpleNamespace( + global_batch=SimpleNamespace(get_sequence=lambda uuid: seq), + _uuid_to_local_map={}, + _local_to_uuid_map={}, + query_book={}, + _free_local_indices=set(), + rank=0, + _response_queue=response_queue, + _get_finish_reason=lambda _seq: "stop", + ) + + BatchGenWorker._report_completion(worker, seq.uuid, gathered_text="answer") + result = response_queue.get() + + assert result["prompt_length"] == 100 + assert result["decoded_length"] == 128 + + +def test_gather_completed_tokens_applies_stop_token_trimming(monkeypatch): + seq = SequenceEntry( + "stopped", + global_idx=9, + prompt_length=10, + max_decode_length=32, + ) + seq.decoded_length = 2 + + class Tokenizer: + def decode(self, token_ids, *, skip_special_tokens): + assert token_ids == [42] + assert skip_special_tokens is True + return "answer" + + worker = object.__new__(BatchGenWorker) + worker._uuid_to_local_map = {seq.uuid: 0} + worker.global_batch = SimpleNamespace(get_sequence=lambda uuid: seq) + worker.query_book = { + 0: SimpleNamespace(decoded_tokens=torch.tensor([[42, 154827]])) + } + worker.world_size = 1 + worker.eos_token_ids = {154820, 154827, 154829} + worker.pad_token_id = 154820 + worker.detokenization_include_special_tokens = False + worker.tokenizer = Tokenizer() + + def gather(output, local): + output[0] = local + + monkeypatch.setattr(worker_module.dist, "all_gather_object", gather) + + assert worker._gather_completed_tokens([seq.uuid]) == {seq.uuid: "answer"} + + +def test_report_completion_fallback_applies_stop_token_trimming(): + seq = SequenceEntry( + "local-fallback", + global_idx=11, + prompt_length=10, + max_decode_length=32, + ) + seq.decoded_tokens = torch.tensor([[42, 154827]]) + seq.decoded_length = 2 + seq.status = SequenceStatus.COMPLETED + + class Tokenizer: + def decode(self, token_ids, *, skip_special_tokens): + assert token_ids == [42] + assert skip_special_tokens is True + return "answer" + + worker = object.__new__(BatchGenWorker) + worker.global_batch = SimpleNamespace(get_sequence=lambda uuid: seq) + worker._uuid_to_local_map = {} + worker._local_to_uuid_map = {} + worker.query_book = {} + worker._free_local_indices = set() + worker.rank = 0 + worker._response_queue = SimpleQueue() + worker.eos_token_ids = {154820, 154827, 154829} + worker.pad_token_id = 154820 + worker.detokenization_include_special_tokens = False + worker.tokenizer = Tokenizer() + worker._get_finish_reason = lambda _seq: "stop" + + worker._report_completion(seq.uuid) + + assert worker._response_queue.get()["text"] == "answer"