diff --git a/serve.py b/serve.py new file mode 100644 index 0000000..2e58c3b --- /dev/null +++ b/serve.py @@ -0,0 +1,97 @@ +#!/usr/bin/env python3 +"""Minimal CRI server. Exposes teach/trigger/generate endpoints. + +Usage: + pip install fastapi uvicorn + python serve.py --model Qwen/Qwen2.5-0.5B --port 8811 + +API: + POST /teach {"prompt": "...", "answer": "..."} + POST /trigger {"query": "...", "max_tokens": 10} + POST /save {"path": "bank.json"} + POST /load {"path": "bank.json"} + GET /stats +""" + +import argparse +import sys +sys.path.insert(0, "python") + +from fastapi import FastAPI +from pydantic import BaseModel +import uvicorn + +from epimem import TransformersBackbone, EpisodicMemory, teach_fact, recall_fact + +app = FastAPI(title="CRI Server") +backbone = None +memory = None + + +class TeachRequest(BaseModel): + prompt: str + answer: str + +class TriggerRequest(BaseModel): + query: str + max_tokens: int = 10 + threshold: float = 0.3 + +class PathRequest(BaseModel): + path: str + + +@app.post("/teach") +def teach(req: TeachRequest): + teach_fact(backbone, memory, req.prompt, req.answer) + return {"status": "conditioned", "total_reflexes": len(memory.episodes)} + + +@app.post("/trigger") +def trigger(req: TriggerRequest): + text, sim, episode = recall_fact(backbone, memory, req.query, req.max_tokens) + return { + "text": text.strip(), + "similarity": round(sim, 4), + "triggered": episode is not None, + "matched_prompt": episode["prompt"] if episode else None, + } + + +@app.post("/save") +def save(req: PathRequest): + memory.save(req.path) + return {"status": "saved", "path": req.path} + + +@app.post("/load") +def load(req: PathRequest): + memory.load(req.path) + return {"status": "loaded", "episodes": len(memory.episodes)} + + +@app.get("/stats") +def stats(): + return { + "model": backbone.model_name, + "hidden_dim": backbone.hidden_dim, + "vocab_size": backbone.vocab_size, + "reflexes": len(memory.episodes), + } + + +if __name__ == "__main__": + parser = argparse.ArgumentParser() + parser.add_argument("--model", default="Qwen/Qwen2.5-0.5B") + parser.add_argument("--port", type=int, default=8811) + parser.add_argument("--bank", type=str, default=None, help="Load reflex bank on startup") + args = parser.parse_args() + + print(f"Starting CRI server on port {args.port}...") + backbone = TransformersBackbone(args.model) + memory = EpisodicMemory() + + if args.bank: + memory.load(args.bank) + + uvicorn.run(app, host="0.0.0.0", port=args.port)