98 lines
2.4 KiB
Python
98 lines
2.4 KiB
Python
#!/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)
|