Skip to content
Draft
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
16 changes: 11 additions & 5 deletions batchgen/batchgen_worker.py
Original file line number Diff line number Diff line change
Expand Up @@ -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({
Expand All @@ -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),
})
Expand All @@ -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
Expand Down
109 changes: 109 additions & 0 deletions tests/test_completion_usage_accounting.py
Original file line number Diff line number Diff line change
@@ -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"
Loading