ares-static-lab / ares_core /api_server.py
jacmor64's picture
Add Ares checkpoint API connector for static UI
479272c verified
Raw
History Blame Contribute Delete
4.91 kB
from __future__ import annotations
import argparse
from dataclasses import dataclass
from typing import Optional
def choose_device(requested: str):
import torch
if requested == "auto":
return "cuda" if torch.cuda.is_available() else "cpu"
return requested
@dataclass
class RuntimeState:
model: object
tokenizer: object
device: str
eos_id: Optional[int]
STATE: Optional[RuntimeState] = None
def load_runtime(checkpoint: str, tokenizer_path: str, device: str = "auto") -> RuntimeState:
import torch
from tokenizers import Tokenizer
from .config import AresConfig
from .model import AresForCausalLM
resolved_device = choose_device(device)
ckpt = torch.load(checkpoint, map_location=resolved_device)
cfg = AresConfig(**ckpt["config"])
model = AresForCausalLM(cfg).to(resolved_device)
state = {k.replace("_orig_mod.", ""): v for k, v in ckpt["model"].items()}
model.load_state_dict(state, strict=True)
model.eval()
tok = Tokenizer.from_file(tokenizer_path)
eos_id = tok.token_to_id("<|eos|>")
return RuntimeState(model=model, tokenizer=tok, device=resolved_device, eos_id=eos_id)
def create_app(checkpoint: str, tokenizer_path: str, device: str = "auto"):
try:
from fastapi import FastAPI
from fastapi.middleware.cors import CORSMiddleware
from pydantic import BaseModel
except ImportError as exc:
raise SystemExit("Install API dependencies: pip install fastapi uvicorn pydantic") from exc
import torch
global STATE
STATE = load_runtime(checkpoint, tokenizer_path, device=device)
class GenerateRequest(BaseModel):
prompt: str
max_new_tokens: int = 180
temperature: float = 0.75
top_k: int = 50
system_prompt: str = "You are Ares, a from-scratch AI assistant. Be honest, useful, and concise."
chat_format: bool = True
app = FastAPI(title="Ares Checkpoint API", version="0.1.0")
app.add_middleware(
CORSMiddleware,
allow_origins=["*"],
allow_credentials=False,
allow_methods=["*"],
allow_headers=["*"],
)
@app.get("/health")
def health():
assert STATE is not None
cfg = STATE.model.cfg
return {
"ok": True,
"device": STATE.device,
"model_name": cfg.model_name,
"max_seq_len": cfg.max_seq_len,
"vocab_size": cfg.vocab_size,
"n_layers": cfg.n_layers,
"d_model": cfg.d_model,
}
@app.post("/generate")
def generate(req: GenerateRequest):
assert STATE is not None
tok = STATE.tokenizer
model = STATE.model
prompt = req.prompt.strip()
if req.chat_format:
prompt_text = (
f"<|system|>\n{req.system_prompt}\n<|end|>\n"
f"<|user|>\n{prompt}\n<|end|>\n"
f"<|assistant|>\n"
)
else:
prompt_text = prompt
enc = tok.encode(prompt_text)
max_input = max(1, model.cfg.max_seq_len - max(1, req.max_new_tokens) - 1)
ids = enc.ids[-max_input:]
x = torch.tensor(ids, dtype=torch.long, device=STATE.device)[None, :]
with torch.no_grad():
out = model.generate(
x,
max_new_tokens=max(1, min(int(req.max_new_tokens), model.cfg.max_seq_len - x.size(1))),
temperature=float(req.temperature),
top_k=int(req.top_k),
eos_id=STATE.eos_id,
)
text = tok.decode(out[0].tolist())
answer = text
marker = "<|assistant|>"
if marker in answer:
answer = answer.split(marker)[-1]
# Remove trailing special markers best-effort.
for stop in ["<|eos|>", "<|end|>", "<|user|>", "<|system|>"]:
if stop in answer:
answer = answer.split(stop)[0]
return {
"text": answer.strip(),
"full_text": text,
"model_name": model.cfg.model_name,
"device": STATE.device,
}
return app
def main() -> None:
parser = argparse.ArgumentParser(description="Serve an Ares checkpoint through a small HTTP API.")
parser.add_argument("--checkpoint", required=True)
parser.add_argument("--tokenizer", required=True)
parser.add_argument("--device", default="auto")
parser.add_argument("--host", default="0.0.0.0")
parser.add_argument("--port", type=int, default=8000)
args = parser.parse_args()
try:
import uvicorn
except ImportError as exc:
raise SystemExit("Install API dependencies: pip install fastapi uvicorn pydantic") from exc
app = create_app(args.checkpoint, args.tokenizer, device=args.device)
uvicorn.run(app, host=args.host, port=args.port)
if __name__ == "__main__":
main()