| """ |
| Cerebellum-2B High-Performance Decision Server |
| Includes REST API and Built-in Interactive Web UI |
| """ |
| import os |
| import sys |
| import time |
| from typing import List, Dict, Optional |
| import torch |
| import uvicorn |
| from fastapi import FastAPI, HTTPException |
| from fastapi.responses import HTMLResponse |
| from pydantic import BaseModel, Field |
|
|
| |
| sys.path.append(os.path.dirname(os.path.abspath(__file__))) |
| from modeling_cerebellum import CerebellumModel |
|
|
| app = FastAPI( |
| title="Cerebellum-2B Decision Server", |
| description="25ms Non-Autoregressive Agent System 1 Decision Engine", |
| version="1.0.0" |
| ) |
|
|
| |
| model: Optional[CerebellumModel] = None |
|
|
| class DecideRequest(BaseModel): |
| state: str = Field(..., description="Agent dialogue history, environment observation, or context") |
| candidates: List[str] = Field(..., min_items=1, description="List of candidate API calls, tools, or DOM actions") |
| instruction: Optional[str] = Field("Select the best action to execute next.", description="Decision instruction") |
| escalate_threshold: Optional[float] = Field(0.50, description="Escalate to human/LLM threshold") |
|
|
| class DecideResponse(BaseModel): |
| action: str |
| action_index: int |
| confidence: float |
| probabilities: Dict[str, float] |
| needs_escalation: bool |
| escalate_probability: float |
| latency_ms: float |
|
|
| class BatchDecideRequest(BaseModel): |
| queries: List[DecideRequest] |
|
|
| class BatchDecideResponse(BaseModel): |
| results: List[DecideResponse] |
| total_latency_ms: float |
|
|
| @app.on_event("startup") |
| def startup(): |
| global model |
| model_dir = os.path.dirname(os.path.abspath(__file__)) |
| device = "cuda:0" if torch.cuda.is_available() else "cpu" |
| print(f"[Cerebellum] Loading model from {model_dir} on {device}...") |
| model = CerebellumModel.from_pretrained(model_dir, device=device) |
| |
| _ = model.decide("Hello", ["Action A", "Action B"]) |
| print("[Cerebellum] Model ready for fast non-autoregressive decisions!") |
|
|
| @app.get("/health") |
| def health(): |
| return {"status": "ok", "model": "Cerebellum-2B", "device": "cuda:0" if torch.cuda.is_available() else "cpu"} |
|
|
| @app.post("/v1/decide", response_model=DecideResponse) |
| def decide(req: DecideRequest): |
| if model is None: |
| raise HTTPException(status_code=503, detail="Model is still initializing") |
| if len(req.candidates) == 0: |
| raise HTTPException(status_code=400, detail="Candidates list cannot be empty") |
| |
| t0 = time.perf_counter() |
| dec = model.decide( |
| state=req.state, |
| candidates=req.candidates, |
| instruction=req.instruction, |
| escalate_threshold=req.escalate_threshold |
| ) |
| return DecideResponse( |
| action=dec.action, |
| action_index=dec.action_index, |
| confidence=dec.confidence, |
| probabilities=dec.probabilities, |
| needs_escalation=dec.needs_escalation, |
| escalate_probability=dec.escalate_probability, |
| latency_ms=dec.latency_ms |
| ) |
|
|
| @app.post("/v1/batch_decide", response_model=BatchDecideResponse) |
| def batch_decide(req: BatchDecideRequest): |
| if model is None: |
| raise HTTPException(status_code=503, detail="Model is still initializing") |
| t0 = time.perf_counter() |
| results = [] |
| for q in req.queries: |
| dec = model.decide( |
| state=q.state, |
| candidates=q.candidates, |
| instruction=q.instruction, |
| escalate_threshold=q.escalate_threshold |
| ) |
| results.append(DecideResponse( |
| action=dec.action, |
| action_index=dec.action_index, |
| confidence=dec.confidence, |
| probabilities=dec.probabilities, |
| needs_escalation=dec.needs_escalation, |
| escalate_probability=dec.escalate_probability, |
| latency_ms=dec.latency_ms |
| )) |
| total_latency = (time.perf_counter() - t0) * 1000 |
| return BatchDecideResponse(results=results, total_latency_ms=total_latency) |
|
|
| @app.get("/", response_class=HTMLResponse) |
| def dashboard(): |
| return """<!DOCTYPE html> |
| <html lang="en"> |
| <head> |
| <meta charset="UTF-8"> |
| <meta name="viewport" content="width=device-width, initial-scale=1.0"> |
| <title>Cerebellum-2B Interactive Decision Console</title> |
| <link href="https://cdn.jsdelivr.net/npm/bootstrap@5.3.0/dist/css/bootstrap.min.css" rel="stylesheet"> |
| <style> |
| body { background: #0f172a; color: #f8fafc; font-family: -apple-system, BlinkMacSystemFont, "Segoe UI", Roboto, sans-serif; padding-top: 2rem; } |
| .card { background: #1e293b; border: 1px solid #334155; border-radius: 12px; } |
| .btn-cerebellum { background: linear-gradient(135deg, #6366f1, #8b5cf6); color: white; font-weight: 600; border: none; } |
| .btn-cerebellum:hover { background: linear-gradient(135deg, #4f46e5, #7c3aed); color: white; } |
| .badge-fast { background: #10b981; color: white; font-size: 0.9rem; padding: 0.4rem 0.8rem; border-radius: 20px; } |
| .badge-esc { background: #ef4444; color: white; font-size: 0.9rem; padding: 0.4rem 0.8rem; border-radius: 20px; } |
| .prob-bar { height: 26px; border-radius: 6px; background: #334155; overflow: hidden; margin-bottom: 8px; } |
| .prob-fill { height: 100%; background: linear-gradient(90deg, #6366f1, #a855f7); transition: width 0.4s ease; display: flex; align-items: center; padding-left: 10px; font-weight: bold; font-size: 0.85rem; color: white; } |
| textarea, input { background: #0f172a !important; color: #f8fafc !important; border: 1px solid #334155 !important; } |
| </style> |
| </head> |
| <body> |
| <div class="container" style="max-width: 900px;"> |
| <div class="d-flex align-items-center justify-content-between mb-4"> |
| <div> |
| <h2 class="fw-bold mb-1">🧠 Cerebellum-2B (小脑-2B)</h2> |
| <p class="text-secondary mb-0">25ms Non-Autoregressive AI Agent System 1 Decision Engine</p> |
| </div> |
| <span class="badge-fast">⚡ O(1) Single Forward Pass</span> |
| </div> |
| |
| <div class="card p-4 shadow-lg mb-4"> |
| <div class="mb-3"> |
| <label class="form-label fw-semibold">Agent State (Environment Observation / Context):</label> |
| <textarea id="state" class="form-control" rows="4">User: My flight was cancelled due to weather. I need to rebook to the earliest flight tomorrow or get a full refund. |
| Flight: CA1832, PNR: X8J29A |
| Status: Flight marked cancelled in airline system.</textarea> |
| </div> |
| |
| <div class="mb-3"> |
| <label class="form-label fw-semibold">Candidate Tools / Actions (One per line):</label> |
| <textarea id="candidates" class="form-control" rows="4">Tool: search_rebooking_options(pnr='X8J29A', date='tomorrow', max_options=3) |
| Tool: issue_involuntary_refund(pnr='X8J29A', reason='weather_cancellation') |
| Tool: charge_rebooking_fee(pnr='X8J29A', amount=50) |
| Tool: escalate_to_human_supervisor(reason='weather_mass_disruption')</textarea> |
| </div> |
| |
| <button class="btn btn-cerebellum py-2 w-100" onclick="makeDecision()">🚀 Execute O(1) Fast Decision</button> |
| </div> |
| |
| <div id="result-card" class="card p-4 shadow-lg d-none"> |
| <div class="d-flex align-items-center justify-content-between mb-3"> |
| <h4 class="fw-bold mb-0">Decision Result</h4> |
| <div id="badges"></div> |
| </div> |
| |
| <div class="alert alert-dark border border-secondary mb-3" id="best-action-box"> |
| <div class="text-secondary small">Selected Action:</div> |
| <div class="fw-bold fs-5 text-warning" id="best-action"></div> |
| </div> |
| |
| <h6 class="fw-semibold mb-2">Candidate Probability Distribution:</h6> |
| <div id="prob-container"></div> |
| </div> |
| </div> |
| |
| <script> |
| async function makeDecision() { |
| const state = document.getElementById('state').value.trim(); |
| const cands = document.getElementById('candidates').value.trim().split('\n').map(s => s.trim()).filter(s => s.length > 0); |
| if (!state || cands.length === 0) return alert('Please provide state and at least 1 candidate action'); |
| |
| const res = await fetch('/v1/decide', { |
| method: 'POST', |
| headers: { 'Content-Type': 'application/json' }, |
| body: JSON.stringify({ state, candidates: cands }) |
| }); |
| const data = await res.json(); |
| |
| document.getElementById('result-card').classList.remove('d-none'); |
| document.getElementById('best-action').innerText = data.action; |
| |
| const badges = document.getElementById('badges'); |
| badges.innerHTML = ` |
| <span class="badge bg-success me-2">${data.latency_ms.toFixed(1)} ms</span> |
| <span class="badge bg-primary me-2">Confidence: ${(data.confidence * 100).toFixed(1)}%</span> |
| <span class="badge ${data.needs_escalation ? 'bg-danger' : 'bg-secondary'}"> |
| ${data.needs_escalation ? '⚠️ Escalate to Human/LLM' : '✅ Autonomous Execute'} |
| </span> |
| `; |
| |
| const container = document.getElementById('prob-container'); |
| container.innerHTML = ''; |
| for (const [action, p] of Object.entries(data.probabilities)) { |
| const pct = (p * 100).toFixed(1); |
| container.innerHTML += ` |
| <div class="small mb-1 text-light d-flex justify-content-between"> |
| <span class="text-truncate" style="max-width: 80%;">${action}</span> |
| <span class="fw-bold">${pct}%</span> |
| </div> |
| <div class="prob-bar"> |
| <div class="prob-fill" style="width: ${pct}%;"></div> |
| </div> |
| `; |
| } |
| } |
| </script> |
| </body> |
| </html>""" |
|
|
| if __name__ == "__main__": |
| uvicorn.run("serve:app", host="0.0.0.0", port=8000, workers=1) |
|
|