Logprobs Refractor (#331)

This commit is contained in:
Liangsheng Yin
2024-03-28 14:34:49 +08:00
committed by GitHub
parent 24e59f5350
commit 3842eba5fa
14 changed files with 385 additions and 152 deletions
+122 -35
View File
@@ -123,31 +123,97 @@ async def flush_cache():
)
async def detokenize_logprob_tokens(token_logprobs):
token_ids = [tid for tid, _ in token_logprobs]
async def detokenize_logprob_tokens(token_logprobs, decode_to_text):
if not decode_to_text:
return [(logprob, token_id, None) for logprob, token_id in token_logprobs]
token_ids = [tid for _, tid in token_logprobs]
token_texts = await tokenizer_manager.detokenize(DetokenizeReqInput(token_ids))
return [(text, logprob) for text, (_, logprob) in zip(token_texts, token_logprobs)]
return [
(logprob, token_id, token_text)
for (logprob, token_id), token_text, in zip(token_logprobs, token_texts)
]
async def detokenize_top_logprobs_tokens(top_logprobs, decode_to_text):
for i, t in enumerate(top_logprobs):
if top_logprobs[i] is not None:
top_logprobs[i] = await detokenize_logprob_tokens(t, decode_to_text)
return top_logprobs
async def handle_token_logprobs_results(obj: GenerateReqInput, ret):
"""Handle the token logprobs results, convert token ids to text if needed.
Args:
obj (GenerateReqInput): The request object.
ret (Union[Dict, List[Dict]]): The response object.
"""
# NOTE: This is because the multiple requests in one http request.
async def convert_style(r, return_text):
r["meta_info"]["prefill_token_logprobs"] = await detokenize_logprob_tokens(
r["meta_info"]["prefill_token_logprobs"], return_text
)
r["meta_info"]["decode_token_logprobs"] = await detokenize_logprob_tokens(
r["meta_info"]["decode_token_logprobs"], return_text
)
r["meta_info"]["prefill_top_logprobs"] = await detokenize_top_logprobs_tokens(
r["meta_info"]["prefill_top_logprobs"], return_text
)
r["meta_info"]["decode_top_logprobs"] = await detokenize_top_logprobs_tokens(
r["meta_info"]["decode_top_logprobs"], return_text
)
if isinstance(obj.text, str):
if obj.return_logprob:
await convert_style(ret, obj.return_text_in_logprobs)
else:
for i, r in enumerate(ret):
if obj.return_logprob[i]:
await convert_style(r, obj.return_text_in_logprobs)
async def stream_generator(obj: GenerateReqInput):
async for out in tokenizer_manager.generate_request(obj):
if obj.return_logprob and obj.return_text_in_logprobs:
out["meta_info"]["token_logprob"] = await detokenize_logprob_tokens(
out["meta_info"]["token_logprob"]
)
await handle_token_logprobs_results(obj, out)
yield out
async def make_openai_style_logprobs(token_logprobs):
async def make_openai_style_logprobs(
prefill_token_logprobs=None,
decode_token_logprobs=None,
prefill_top_logprobs=None,
decode_top_logprobs=None,
):
ret_logprobs = LogProbs()
for token_text, token_logprob in token_logprobs:
ret_logprobs.tokens.append(token_text)
ret_logprobs.token_logprobs.append(token_logprob)
def append_token_logprobs(token_logprobs):
for logprob, _, token_text in token_logprobs:
ret_logprobs.tokens.append(token_text)
ret_logprobs.token_logprobs.append(logprob)
# Not Supported yet
ret_logprobs.text_offset.append(-1)
def append_top_logprobs(top_logprobs):
for tokens in top_logprobs:
if tokens is not None:
ret_logprobs.top_logprobs.append(
{token[2]: token[0] for token in tokens}
)
else:
ret_logprobs.top_logprobs.append(None)
if prefill_token_logprobs is not None:
append_token_logprobs(prefill_token_logprobs)
if decode_token_logprobs is not None:
append_token_logprobs(decode_token_logprobs)
if prefill_top_logprobs is not None:
append_top_logprobs(prefill_top_logprobs)
if decode_top_logprobs is not None:
append_top_logprobs(decode_top_logprobs)
# Not supported yet.
ret_logprobs.top_logprobs.append({})
ret_logprobs.text_offset.append(-1)
return ret_logprobs
@@ -165,10 +231,7 @@ async def generate_request(obj: GenerateReqInput):
return StreamingResponse(stream_results(), media_type="text/event-stream")
ret = await tokenizer_manager.generate_request(obj).__anext__()
if obj.return_logprob and obj.return_text_in_logprobs:
ret["meta_info"]["token_logprob"] = await detokenize_logprob_tokens(
ret["meta_info"]["token_logprob"]
)
await handle_token_logprobs_results(obj, ret)
return ret
@@ -192,7 +255,8 @@ async def v1_completions(raw_request: Request):
"frequency_penalty": request.frequency_penalty,
"regex": request.regex,
},
return_logprob=request.logprobs is not None,
return_logprob=request.logprobs is not None and request.logprobs > 0,
top_logprobs_num=request.logprobs if request.logprobs is not None else 0,
return_text_in_logprobs=True,
stream=request.stream,
)
@@ -212,15 +276,32 @@ async def v1_completions(raw_request: Request):
if request.echo:
# Prepend prompt in response text.
text = request.prompt + text
else:
# Skip prompt tokens if echo is disabled.
n_prev_token = prompt_tokens
if request.logprobs is not None:
if request.logprobs:
# The first chunk and echo is enabled.
if not stream_buffer and request.echo:
prefill_token_logprobs = content["meta_info"][
"prefill_token_logprobs"
]
prefill_top_logprobs = content["meta_info"][
"prefill_top_logprobs"
]
else:
prefill_token_logprobs = None
prefill_top_logprobs = None
logprobs = await make_openai_style_logprobs(
content["meta_info"]["token_logprob"][n_prev_token:]
prefill_token_logprobs=prefill_token_logprobs,
prefill_top_logprobs=prefill_top_logprobs,
decode_token_logprobs=content["meta_info"][
"decode_token_logprobs"
][n_prev_token:],
decode_top_logprobs=content["meta_info"]["decode_top_logprobs"][
n_prev_token:
],
)
n_prev_token = len(content["meta_info"]["token_logprob"])
n_prev_token = len(content["meta_info"]["decode_token_logprobs"])
else:
logprobs = None
@@ -255,20 +336,26 @@ async def v1_completions(raw_request: Request):
prompt_tokens = ret["meta_info"]["prompt_tokens"]
completion_tokens = ret["meta_info"]["completion_tokens"]
text = ret["text"]
token_logprob_pos = prompt_tokens
if request.echo:
token_logprob_pos = 0
text = request.prompt + text
else:
token_logprob_pos = prompt_tokens
logprobs = (
await make_openai_style_logprobs(
ret["meta_info"]["token_logprob"][token_logprob_pos:]
if request.logprobs:
if request.echo:
prefill_token_logprobs = ret["meta_info"]["prefill_token_logprobs"]
prefill_top_logprobs = ret["meta_info"]["prefill_top_logprobs"]
else:
prefill_token_logprobs = None
prefill_top_logprobs = None
logprobs = await make_openai_style_logprobs(
prefill_token_logprobs=prefill_token_logprobs,
prefill_top_logprobs=prefill_top_logprobs,
decode_token_logprobs=ret["meta_info"]["decode_token_logprobs"],
decode_top_logprobs=ret["meta_info"]["decode_top_logprobs"],
)
if request.logprobs is not None
else None
)
else:
logprobs = None
choice_data = CompletionResponseChoice(
index=0,
text=text,