Support decode token logprobs (#130)
This commit is contained in:
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user