prompt streaming
This commit is contained in:
@@ -1,7 +1,33 @@
|
||||
#python << '#EOF'
|
||||
# vim: ft=python ts=4 sw=4 et :
|
||||
import requests
|
||||
import os, sys, json, pathlib, argparse
|
||||
import asyncio, enum, os, sys, json, pathlib, argparse
|
||||
import httpx
|
||||
|
||||
from rich import print, print_json
|
||||
from rich.console import Console
|
||||
from rich.prompt import Prompt
|
||||
from rich.json import JSON
|
||||
from rich.markdown import Markdown
|
||||
|
||||
output={sys.stdout: Console()}
|
||||
class bb(enum.StrEnum):
|
||||
thinking="bold yellow"
|
||||
content="bold blue"
|
||||
model="bold red"
|
||||
prompt="bold green"
|
||||
|
||||
def write_rich(*msg, end='\n', stream=sys.stdout):
|
||||
output[stream].print(*msg, end=end)
|
||||
if end != '\n':
|
||||
stream.flush()
|
||||
|
||||
def write_basic(*msg, sep=' ', end='\n', stream=sys.stdout):
|
||||
stream.write(sep.join(*msg) + end)
|
||||
if end != '\n':
|
||||
stream.flush()
|
||||
|
||||
write = write_rich
|
||||
|
||||
def parse_args():
|
||||
p = argparse.ArgumentParser()
|
||||
@@ -13,60 +39,111 @@ def parse_args():
|
||||
"llama3.1:8b",
|
||||
"qwen3.5:4b",
|
||||
])
|
||||
p.add_argument("--prompt")
|
||||
p.add_argument("--system-prompt",
|
||||
default="")
|
||||
|
||||
g = p.add_mutually_exclusive_group()
|
||||
g.add_argument("--pull")
|
||||
g.add_argument("--chat", action="store_true")
|
||||
g.add_argument("--generate", action="store_true")
|
||||
g.add_argument("--endpoint",
|
||||
mxg = p.add_mutually_exclusive_group()
|
||||
mxg.add_argument("--pull")
|
||||
#mxg.add_argument("--chat", action=argparse.BooleanOptionalAction, default=True)
|
||||
mxg.add_argument("--response", action="store_true")
|
||||
mxg.add_argument("--endpoint",
|
||||
default="/api/chat",
|
||||
choices=[
|
||||
"/api/chat",
|
||||
"/v1/chat/completions",])
|
||||
g.add_argument("--api",
|
||||
mxg.add_argument("--api",
|
||||
default="ollama",
|
||||
choices=["ollama","openai",])
|
||||
|
||||
url = p.add_mutually_exclusive_group()
|
||||
url.add_argument("--url", default=os.getenv("PROMPT_URL", "http://localhost:11434"))
|
||||
url.add_argument("--host")
|
||||
url.add_argument("--host-port")
|
||||
url.add_argument("--port")
|
||||
g = p.add_argument_group("Model Options")
|
||||
g.add_argument("--reasoning", action="store_true")
|
||||
g.add_argument("--stream", action="store_true")
|
||||
g.add_argument("--num_ctx", default=128000, type=int)
|
||||
g.add_argument("--top_k", default=40, type=int)
|
||||
g.add_argument("--top_p", default=0.9, type=float)
|
||||
g.add_argument("--temperature", default=0.2, type=float)
|
||||
|
||||
url = p.add_argument_group("URL")
|
||||
url.add_argument("--url",
|
||||
default=os.getenv("PROMPT_URL", "http://localhost:11434"))
|
||||
url.add_argument("--host",
|
||||
default=os.getenv("PROMPT_HOST", "localhost"))
|
||||
url.add_argument("--port",
|
||||
default=os.getenv("PROMPT_PORT", "11434"))
|
||||
url.add_argument("--proto",
|
||||
default=os.getenv("PROMPT_PROTO", "http"))
|
||||
url.add_argument("--apikey")
|
||||
url.add_argument("--timeout", type=int, default=180)
|
||||
return p.parse_args()
|
||||
|
||||
|
||||
async def do_async_turn(conf):
|
||||
"""Do user->assistant turn streaming responses"""
|
||||
if not conf.prompt:
|
||||
prompt = Prompt.ask(f"[{bb.model}]Ask {conf.model}»[/] [{bb.prompt}]")
|
||||
if conf.response:
|
||||
system_and_user = {
|
||||
"input": conf.prompt or prompt,
|
||||
"instructions": conf.system_prompt}
|
||||
else:
|
||||
system_and_user = {"messages": [
|
||||
{"role":"system", "content": conf.system_prompt},
|
||||
{"role": "user", "content": conf.prompt or prompt}]}
|
||||
|
||||
client = httpx.AsyncClient(
|
||||
base_url="http://192.168.14.20:11434",
|
||||
headers={
|
||||
"Authorization": f"Bearer {conf.apikey}",},
|
||||
timeout=float(conf.timeout))
|
||||
async with client.stream('POST',"/api/chat", json=system_and_user | {
|
||||
"model": conf.model,
|
||||
"reasoning": {"enabled": conf.reasoning},
|
||||
"stream": conf.stream,
|
||||
"options": {
|
||||
"num_ctx": conf.num_ctx,
|
||||
"top_k": conf.top_k,
|
||||
"top_p": conf.top_p,
|
||||
"temperature": conf.temperature,
|
||||
}}) as res:
|
||||
msg = {"thinking": {"len": 1, 1: []},
|
||||
"content": {"len": 1, 1: []}}
|
||||
async for line in res.aiter_lines():
|
||||
if not line:
|
||||
write('[i grey].[/]', end='')
|
||||
# TODO: use spinner from rich
|
||||
continue
|
||||
data = json.loads(line)
|
||||
if data.get("done", False):
|
||||
write(data)
|
||||
msg['done'] = data
|
||||
return msg
|
||||
#write(f"[i]{data}[/i]")
|
||||
for i in ["thinking","content"]:
|
||||
text = data.get("message", {}).get(i, "")
|
||||
msg[i][msg[i]['len']].append(text)
|
||||
write(f"[i]{text}[/i]", end='')
|
||||
if '\\n' in msg[i][msg[i]['len']]:
|
||||
write(f"[{bb.line}]{msg[i]['len']:>3d}:[/] "
|
||||
f"[{bb(i)}]{Markdown(msg[i][msg[i]['len']])}[/]")
|
||||
msg[i]['len'] += 1
|
||||
|
||||
def main():
|
||||
conf = parse_args()
|
||||
model_slug = conf.model
|
||||
|
||||
if conf.prompt:
|
||||
write(f"[{bb.model}]Ask {conf.model}»[/] [{bb.prompt}]{conf.prompt}[/]")
|
||||
try:
|
||||
#while True:
|
||||
#prompt = input(model_slug + "> ")
|
||||
prompt = "What is the meaning of life?"
|
||||
res = requests.post(
|
||||
url="http://192.168.14.20:11434/api/chat",
|
||||
#headers={
|
||||
#"Authorization": f"Bearer {OPENROUTER_API_KEY}",
|
||||
#"HTTP-Referer": "<YOUR_SITE_URL>", # Optional. Site URL for rankings on openrouter.ai.
|
||||
#"X-OpenRouter-Title": "<YOUR_SITE_NAME>", # Optional. Site title for rankings on openrouter.ai.
|
||||
#},
|
||||
data=json.dumps({
|
||||
"model": model_slug,
|
||||
"messages": [{"role": "user","content": prompt}],
|
||||
"reasoning": {"enabled": True},
|
||||
"stream": False,
|
||||
"options": {
|
||||
"num_ctx": 128000,
|
||||
"top_k": 40,
|
||||
"top_p": 0.2,
|
||||
"temperature": 0.1,
|
||||
}
|
||||
})
|
||||
)
|
||||
|
||||
data = res.json()
|
||||
print(f"{data['message']['role']}: {data['message']['content']}")
|
||||
#response = response['choices'][0]['message']
|
||||
while True:
|
||||
if conf.stream:
|
||||
data = asyncio.run(do_async_turn(conf))
|
||||
else:
|
||||
data = res.json()
|
||||
if "thinking" in data.get("message",{}):
|
||||
write(f"[{bb.model}]{conf.model}'s thinking:[/] [{bb.thinking}] {Markdown(data['message']['thinking'])}[/]")
|
||||
write(f"[{bb.model}]{conf.model}'s response: [/][{bb.content}]{Markdown(data['message']['content'])}[/]")
|
||||
#response = response['choices'][0]['message']
|
||||
if conf.prompt:
|
||||
sys.exit(0)
|
||||
# TODO: add turn to message list and do another turn
|
||||
except KeyboardInterrupt as e:
|
||||
sys.exit(1)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user