Spaces:
Running on Zero
Running on Zero
File size: 4,733 Bytes
a0270e2 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 | """
FastAPI Server for Parallel Constrained Decision Engine.
Serves interactive side-by-side benchmark UI, presets, and live streaming endpoints.
"""
import os
import json
import asyncio
from typing import Dict, Any, Optional
from fastapi import FastAPI, HTTPException
from fastapi.responses import HTMLResponse, StreamingResponse
from fastapi.staticfiles import StaticFiles
from fastapi.middleware.cors import CORSMiddleware
from pydantic import BaseModel, Field
from core.schema import StructuredSchema
from core.engine import (
get_engine,
run_naive_generation,
stream_naive_generation,
run_parallel_generation,
run_rlcd_generation,
)
app = FastAPI(title="Parallel Constrained Decision Engine")
app.add_middleware(
CORSMiddleware,
allow_origins=["*"],
allow_credentials=True,
allow_methods=["*"],
allow_headers=["*"],
)
PRESETS_DIR = os.path.join(os.path.dirname(os.path.dirname(__file__)), "presets")
WEB_DIR = os.path.join(os.path.dirname(os.path.dirname(__file__)), "web")
class PredictRequest(BaseModel):
context: str
schema_def: Dict[str, Any] = Field(..., alias="schema")
temperature: Optional[float] = None
class Config:
populate_by_name = True
@app.on_event("startup")
def on_startup():
print("Pre-warming inference engine on Apple Silicon GPU...")
get_engine()
print("Engine ready for high-speed inference.")
@app.get("/api/presets")
def list_presets():
presets = []
if os.path.exists(PRESETS_DIR):
for fname in sorted(os.listdir(PRESETS_DIR)):
if fname.endswith(".json"):
fpath = os.path.join(PRESETS_DIR, fname)
try:
with open(fpath, "r") as f:
presets.append(json.load(f))
except Exception as e:
print(f"Error loading preset {fname}: {e}")
return presets
@app.post("/api/run-parallel")
@app.post("/api/run-rlcd")
def api_run_parallel(req: PredictRequest):
try:
schema = StructuredSchema(req.schema_def)
temp = req.temperature if req.temperature is not None else 1.0
res = run_parallel_generation(req.context, schema, temperature=temp)
return res
except Exception as e:
raise HTTPException(status_code=400, detail=str(e))
@app.post("/api/run-naive")
def api_run_naive(req: PredictRequest):
try:
schema = StructuredSchema(req.schema_def)
temp = req.temperature if req.temperature is not None else 0.2
res = run_naive_generation(req.context, schema, temperature=temp)
return res
except Exception as e:
raise HTTPException(status_code=400, detail=str(e))
@app.post("/api/stream-naive")
def api_stream_naive(req: PredictRequest):
"""Server-Sent Events endpoint streaming individual tokens as they are decoded."""
try:
schema = StructuredSchema(req.schema_def)
temp = req.temperature if req.temperature is not None else 0.2
def event_generator():
try:
for event in stream_naive_generation(req.context, schema, temperature=temp):
yield f"data: {json.dumps(event)}\n\n"
except Exception as e:
print(f"Error in stream_naive_generation: {e}")
yield f"data: {json.dumps({'type': 'error', 'error': str(e)})}\n\n"
return StreamingResponse(event_generator(), media_type="text/event-stream")
except Exception as e:
raise HTTPException(status_code=400, detail=str(e))
@app.post("/api/compare")
def api_compare(req: PredictRequest):
try:
schema = StructuredSchema(req.schema_def)
naive_temp = req.temperature if req.temperature is not None else 0.2
rlcd_temp = req.temperature if req.temperature is not None else 1.0
# Run naive
naive_res = run_naive_generation(req.context, schema, temperature=naive_temp)
# Run RLCD
rlcd_res = run_rlcd_generation(req.context, schema, temperature=rlcd_temp)
speedup = naive_res["elapsed_ms"] / max(rlcd_res["elapsed_ms"], 1.0)
steps_reduction = naive_res["sequential_forward_passes"] / max(rlcd_res["sequential_forward_passes"], 1.0)
return {
"speedup_multiplier": round(speedup, 1),
"steps_reduction": round(steps_reduction, 1),
"naive": naive_res,
"parallel": rlcd_res,
"rlcd": rlcd_res
}
except Exception as e:
raise HTTPException(status_code=400, detail=str(e))
# Mount web frontend
if os.path.exists(WEB_DIR):
app.mount("/", StaticFiles(directory=WEB_DIR, html=True), name="static")
|