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()