added token buffer for catchiing stop sequences

This commit is contained in:
Yannic Kilcher
2023-02-09 23:46:44 +01:00
parent aa9b2b2325
commit 4076afd0d8
6 changed files with 126 additions and 21 deletions
+3
View File
@@ -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...")
+23 -21
View File
@@ -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()
)
+56
View File
@@ -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
+1
View File
@@ -1,4 +1,5 @@
loguru
pydantic
rel
requests
sseclient-py
+40
View File
@@ -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
@@ -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):