prompt streaming
This commit is contained in:
@@ -1,7 +1,33 @@
|
|||||||
#python << '#EOF'
|
#python << '#EOF'
|
||||||
# vim: ft=python ts=4 sw=4 et :
|
# vim: ft=python ts=4 sw=4 et :
|
||||||
import requests
|
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():
|
def parse_args():
|
||||||
p = argparse.ArgumentParser()
|
p = argparse.ArgumentParser()
|
||||||
@@ -13,60 +39,111 @@ def parse_args():
|
|||||||
"llama3.1:8b",
|
"llama3.1:8b",
|
||||||
"qwen3.5:4b",
|
"qwen3.5:4b",
|
||||||
])
|
])
|
||||||
|
p.add_argument("--prompt")
|
||||||
|
p.add_argument("--system-prompt",
|
||||||
|
default="")
|
||||||
|
|
||||||
g = p.add_mutually_exclusive_group()
|
mxg = p.add_mutually_exclusive_group()
|
||||||
g.add_argument("--pull")
|
mxg.add_argument("--pull")
|
||||||
g.add_argument("--chat", action="store_true")
|
#mxg.add_argument("--chat", action=argparse.BooleanOptionalAction, default=True)
|
||||||
g.add_argument("--generate", action="store_true")
|
mxg.add_argument("--response", action="store_true")
|
||||||
g.add_argument("--endpoint",
|
mxg.add_argument("--endpoint",
|
||||||
default="/api/chat",
|
default="/api/chat",
|
||||||
choices=[
|
choices=[
|
||||||
"/api/chat",
|
"/api/chat",
|
||||||
"/v1/chat/completions",])
|
"/v1/chat/completions",])
|
||||||
g.add_argument("--api",
|
mxg.add_argument("--api",
|
||||||
default="ollama",
|
default="ollama",
|
||||||
choices=["ollama","openai",])
|
choices=["ollama","openai",])
|
||||||
|
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_mutually_exclusive_group()
|
url = p.add_argument_group("URL")
|
||||||
url.add_argument("--url", default=os.getenv("PROMPT_URL", "http://localhost:11434"))
|
url.add_argument("--url",
|
||||||
url.add_argument("--host")
|
default=os.getenv("PROMPT_URL", "http://localhost:11434"))
|
||||||
url.add_argument("--host-port")
|
url.add_argument("--host",
|
||||||
url.add_argument("--port")
|
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()
|
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():
|
def main():
|
||||||
conf = parse_args()
|
conf = parse_args()
|
||||||
model_slug = conf.model
|
if conf.prompt:
|
||||||
|
write(f"[{bb.model}]Ask {conf.model}»[/] [{bb.prompt}]{conf.prompt}[/]")
|
||||||
try:
|
try:
|
||||||
#while True:
|
while True:
|
||||||
#prompt = input(model_slug + "> ")
|
if conf.stream:
|
||||||
prompt = "What is the meaning of life?"
|
data = asyncio.run(do_async_turn(conf))
|
||||||
res = requests.post(
|
else:
|
||||||
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()
|
data = res.json()
|
||||||
print(f"{data['message']['role']}: {data['message']['content']}")
|
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']
|
#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:
|
except KeyboardInterrupt as e:
|
||||||
sys.exit(1)
|
sys.exit(1)
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user