Team Ai
Apppublic

valeriow/parallel-constrained-decoding

sourceHugging Faceapache-2.0updated 24d agoView on Hugging Face
0likes
app.py143 linesDownload Raw Back to server
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