mirror of
https://github.com/wassname/Open-Assistant.git
synced 2026-07-24 12:50:50 +08:00
added token buffer for catchiing stop sequences
This commit is contained in:
@@ -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...")
|
||||
|
||||
|
||||
@@ -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()
|
||||
)
|
||||
|
||||
@@ -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,4 +1,5 @@
|
||||
loguru
|
||||
pydantic
|
||||
rel
|
||||
requests
|
||||
sseclient-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
|
||||
@@ -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):
|
||||
|
||||
Reference in New Issue
Block a user