diff --git a/inference/text-client/__main__.py b/inference/text-client/__main__.py index 8484978e..e56feaa1 100644 --- a/inference/text-client/__main__.py +++ b/inference/text-client/__main__.py @@ -37,6 +37,9 @@ def main(backend_url: str = "http://127.0.0.1:8000"): data = json.loads(event.data) print(data["token"]["text"], end="", flush=True) print() + except typer.Abort: + typer.echo("Exiting...") + break except Exception: typer.echo("Error, restarting chat...") diff --git a/inference/worker/__main__.py b/inference/worker/__main__.py index cea7f257..2a6514b7 100644 --- a/inference/worker/__main__.py +++ b/inference/worker/__main__.py @@ -1,9 +1,9 @@ -import json - +import interface import rel import requests import sseclient import typer +import utils import websocket from loguru import logger from oasst_shared.schemas import inference, protocol @@ -43,19 +43,12 @@ def main( prompt = prefix + "\n".join(messages) + "\nAssistant:" + parameters = interface.GenerateStreamParameters.from_work_request(work_request) response = requests.post( f"{inference_server_url}/generate_stream", json={ "inputs": prompt, - "parameters": { - "max_new_tokens": work_request.max_new_tokens, - "do_sample": work_request.do_sample, - "top_k": work_request.top_k, - "top_p": work_request.top_p, - "temperature": work_request.temperature, - "seed": work_request.seed, - # "stop": ["\nUser:", "\nAssistant:"], # TODO: make this a bit more workable because it's mutliple tokens - }, + "parameters": parameters.dict(), }, stream=True, headers={"Accept": "text/event-stream"}, @@ -68,26 +61,35 @@ def main( return client = sseclient.SSEClient(response) + stream_response = None + token_buffer = utils.TokenBuffer(stop_sequences=parameters.stop) for event in client.events(): logger.debug(f"Received event: {event}") - data = json.loads(event.data) - if data["generated_text"]: - break - token = data["token"] + stream_response = interface.GenerateStreamResponse.parse_raw(event.data) + token = stream_response.token + for send_token in token_buffer.add(token): + ws.send( + inference.WorkResponsePacket( + token=send_token.to_token_response(), + ).json() + ) + if stream_response is None: + logger.error("No stream response received") + return + + for send_token in token_buffer.finish(reason=stream_response.details.finish_reason): ws.send( inference.WorkResponsePacket( - token=inference.TokenResponse( - text=token["text"], - log_prob=token["logprob"], - token_id=token["id"], - ) + token=send_token.to_token_response(), ).json() ) + ws.send( inference.WorkResponsePacket( is_end=True, generated_text=inference.GeneratedTextResponse( - text=data["generated_text"], + text=stream_response.generated_text, + finish_reason=stream_response.details.finish_reason, ), ).json() ) diff --git a/inference/worker/interface.py b/inference/worker/interface.py new file mode 100644 index 00000000..06a3eac6 --- /dev/null +++ b/inference/worker/interface.py @@ -0,0 +1,56 @@ +from typing import Literal + +import pydantic +from oasst_shared.schemas import inference + + +class GenerateStreamParameters(pydantic.BaseModel): + max_new_tokens: int | None + do_sample: bool | None + top_k: int | None + top_p: float | None + temperature: float | None + repetition_penalty: float | None + seed: int | None + stop: list[str] = ["\nUser:", "\nAssistant:"] # TODO: make this a bit more workable because it's mutliple tokens + details: bool = True + + @staticmethod + def from_work_request(work_request: inference.WorkRequest) -> "GenerateStreamParameters": + return GenerateStreamParameters( + max_new_tokens=work_request.max_new_tokens, + do_sample=work_request.do_sample, + top_k=work_request.top_k, + top_p=work_request.top_p, + temperature=work_request.temperature, + repetition_penalty=work_request.repetition_penalty, + seed=work_request.seed, + ) + + +class Token(pydantic.BaseModel): + text: str + logprob: float + id: int + + def __len__(self) -> int: + return len(self.text) + + def to_token_response(self) -> inference.TokenResponse: + return inference.TokenResponse( + text=self.text, + log_prob=self.logprob, + token_id=self.id, + ) + + +class StreamDetails(pydantic.BaseModel): + generated_tokens: int + seed: int | None + finish_reason: Literal["length", "eos_token", "stop_sequence"] + + +class GenerateStreamResponse(pydantic.BaseModel): + token: Token + generated_text: str | None + details: StreamDetails | None diff --git a/inference/worker/requirements.txt b/inference/worker/requirements.txt index 82169379..3afd3617 100644 --- a/inference/worker/requirements.txt +++ b/inference/worker/requirements.txt @@ -1,4 +1,5 @@ loguru +pydantic rel requests sseclient-py diff --git a/inference/worker/utils.py b/inference/worker/utils.py new file mode 100644 index 00000000..2cababcf --- /dev/null +++ b/inference/worker/utils.py @@ -0,0 +1,40 @@ +import collections +from typing import Literal + +import interface + + +class TokenBuffer: + def __init__(self, stop_sequences: list[str]) -> None: + self.stop_sequences = stop_sequences + self.longest_stop_len = max((len(stop) for stop in stop_sequences), default=0) + self.tokens = collections.deque() + self.token_lens = collections.deque() + self.total_len = 0 + + def add(self, token: interface.Token): + self.tokens.append(token) + self.token_lens.append(len(token)) + self.total_len += len(token) + while True: + if not self.tokens: + break + head_len = self.token_lens[0] + if self.total_len - head_len >= self.longest_stop_len: + token = self.tokens.popleft() + self.token_lens.popleft() + self.total_len -= head_len + yield token + else: + break + + def finish(self, reason: Literal["length", "eos_token", "stop_sequence"]): + if reason == "stop_sequence": + end_sequence = "" + while self.tokens: + end_sequence = self.tokens.pop().text + end_sequence + if end_sequence in self.stop_sequences: + break + yield from self.tokens + else: + yield from self.tokens diff --git a/oasst-shared/oasst_shared/schemas/inference.py b/oasst-shared/oasst_shared/schemas/inference.py index 96b05c7b..1bb89a42 100644 --- a/oasst-shared/oasst_shared/schemas/inference.py +++ b/oasst-shared/oasst_shared/schemas/inference.py @@ -1,4 +1,5 @@ import random +from typing import Literal import pydantic @@ -18,6 +19,7 @@ class WorkRequest(pydantic.BaseModel): top_k: int = 50 top_p: float = 0.9 temperature: float = 1.0 + repetition_penalty: float | None = None class TokenResponse(pydantic.BaseModel): @@ -28,6 +30,7 @@ class TokenResponse(pydantic.BaseModel): class GeneratedTextResponse(pydantic.BaseModel): text: str + finish_reason: Literal["length", "eos_token", "stop_sequence"] class WorkResponsePacket(pydantic.BaseModel):