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