Logprobs Refractor (#331)
This commit is contained in:
+122
-35
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user