Support decode token logprobs (#130)

This commit is contained in:
Cody Yu
2024-02-06 12:24:55 -08:00
committed by GitHub
parent ee1df26a77
commit a7334aeea1
10 changed files with 233 additions and 96 deletions
+34 -14
View File
@@ -388,24 +388,28 @@ class ModelRpcServer(rpyc.Service):
self.model_config.vocab_size, self.int_token_logit_bias
)
logprobs = None
if batch.extend_num_tokens != 0:
# Forward
logits, (logprobs, normalized_logprobs) = self.model_runner.forward(
batch, ForwardMode.EXTEND, batch.return_logprob
logits, (prefill_logprobs, normalized_logprobs, last_logprobs) = (
self.model_runner.forward(batch, ForwardMode.EXTEND, batch.return_logprob)
)
# print("extend logits", logits)
if logprobs is not None:
logprobs = logprobs.cpu().tolist()
if prefill_logprobs is not None:
logprobs = prefill_logprobs.cpu().tolist()
normalized_logprobs = normalized_logprobs.cpu().tolist()
next_token_ids, next_token_probs = batch.sample(logits)
next_token_ids, _ = batch.sample(logits)
next_token_ids = next_token_ids.cpu().tolist()
else:
next_token_ids = [self.tokenizer.eos_token_id] * len(batch.reqs)
logprobs = normalized_logprobs = None
logits = logprobs = normalized_logprobs = last_logprobs = None
# Only batch transfer the selected logprobs of the next token to CPU to reduce overhead.
reqs = batch.reqs
if last_logprobs is not None:
last_logprobs = last_logprobs[torch.arange(len(reqs)), next_token_ids].cpu().tolist()
# Check finish condition
reqs = batch.reqs
pt = 0
for i, req in enumerate(reqs):
req.output_ids = [next_token_ids[i]]
@@ -414,6 +418,10 @@ class ModelRpcServer(rpyc.Service):
if logprobs is not None:
req.logprob = logprobs[pt : pt + req.extend_input_len - 1]
req.normalized_logprob = normalized_logprobs[i]
token_ids = req.input_ids + [next_token_ids[i]]
token_logprobs = [None] + req.logprob + [last_logprobs[i]]
req.token_logprob = list(zip(token_ids, token_logprobs))
pt += req.extend_input_len
self.handle_finished_requests(batch)
@@ -463,15 +471,26 @@ class ModelRpcServer(rpyc.Service):
batch.prepare_for_decode()
# Forward
logits = self.model_runner.forward(batch, ForwardMode.DECODE)
next_token_ids, next_token_probs = batch.sample(logits)
logits, (_, _, last_logprobs) = self.model_runner.forward(
batch,
ForwardMode.DECODE,
batch.return_logprob,
)
next_token_ids, _ = batch.sample(logits)
next_token_ids = next_token_ids.cpu().tolist()
# Check finish condition
# Only batch transfer the selected logprobs of the next token to CPU to reduce overhead.
reqs = batch.reqs
for i in range(len(reqs)):
reqs[i].output_ids.append(next_token_ids[i])
reqs[i].check_finished()
if last_logprobs is not None:
last_logprobs = last_logprobs[torch.arange(len(reqs)), next_token_ids].tolist()
# Check finish condition
for i, (req, next_tok_id) in enumerate(zip(reqs, next_token_ids)):
req.output_ids.append(next_tok_id)
req.check_finished()
if last_logprobs is not None:
req.token_logprob.append((next_tok_id, last_logprobs[i]))
self.handle_finished_requests(batch)
@@ -513,6 +532,7 @@ class ModelRpcServer(rpyc.Service):
}
if req.return_logprob:
meta_info["prompt_logprob"] = req.logprob
meta_info["token_logprob"] = req.token_logprob
meta_info["normalized_prompt_logprob"] = req.normalized_logprob
output_meta_info.append(meta_info)
output_finished.append(req.finished)