Spaces:
Running
Running
| 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 | |
| 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=["*"], | |
| ) | |
| 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, | |
| } | |
| 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() | |