Add minimal FastAPI CRI server (teach/trigger/save/load)
Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
97
serve.py
Normal file
97
serve.py
Normal file
@@ -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)
|
||||
Reference in New Issue
Block a user