"""Rushd-Agent API — Full Pipeline with Early Exit Runs on M2 port 8777 Pipeline: User Query → Router (Qwen2.5:7b) → Expert (Geo) → Early Exit → Output """ import os, sys, json, time, asyncio, subprocess os.environ["WANDB_DISABLED"] = "true" import mlx.core as mx import mlx.nn as nn from mlx_lm import load from mlx_lm.models.qwen3_5 import DecoderLayer from fastapi import FastAPI, HTTPException from pydantic import BaseModel import uvicorn # ─── Config ─── MODEL_PATH = "/Users/ai/rushd-geo-mlx-4bit" PREDICTOR_PATH = "/Users/ai/noema_predictor.safetensors" N_LAYERS = 64 EXIT_LAYER = 6 SKIP_TO = 48 # ─── Load Model ─── print("🚀 Loading Rushd-Geo with Early Exit...", flush=True) t0 = time.time() model, tokenizer = load(MODEL_PATH) print(f" Model loaded in {time.time()-t0:.1f}s", flush=True) # ─── Load Noema Predictor ─── HIDDEN = 5120 class NoemaPredictor(nn.Module): def __init__(self): super().__init__() self.net = nn.Sequential( nn.Linear(HIDDEN, 2048), nn.ReLU(), nn.Linear(2048, HIDDEN), ) def __call__(self, x): return self.net(x) predictor = NoemaPredictor() predictor.load_weights(PREDICTOR_PATH) print(f" Noema Predictor loaded ({PREDICTOR_PATH})", flush=True) # ─── Keep original call ─── original_layer_call = DecoderLayer.__call__ # ─── FastAPI App ─── app = FastAPI(title="Rushd-Agent API", version="1.0.0") class QueryRequest(BaseModel): query: str max_tokens: int = 512 temperature: float = 0.7 expert: str = "auto" # auto, geo, writer, etc. class QueryResponse(BaseModel): output: str expert: str mode: str # early or full nci: float tokens: int elapsed_ms: float saved_layers: int # ─── Router via Ollama ─── EXPERTS = ["geo", "writer", "game", "decision", "plan", "logic"] def route_query(query: str) -> str: """Route query to expert using Qwen2.5:7b router""" result = subprocess.run( ["/opt/homebrew/bin/ollama", "run", "rushd-router-qwen"], input=query.encode(), capture_output=True, timeout=30 ) expert = result.stdout.decode().strip().lower() if expert not in EXPERTS: expert = "geo" return expert # ─── Early Exit Inference ─── def compute_nci(a, b): n_a = a.mean(axis=1).squeeze(0) n_b = b.mean(axis=1).squeeze(0) return float(mx.sum(n_a * n_b) / (mx.linalg.norm(n_a) * mx.linalg.norm(n_b))) def infer_early(input_ids, max_new_tokens=512, temperature=0.7): """Generate with standard inference (full model)""" generated = input_ids.tolist()[0] saved_layers_total = 0 for step in range(max_new_tokens): logits = model(mx.array([generated])) if temperature > 0 and temperature < 0.9: probs = mx.softmax(logits[0, -1, :].astype(mx.float32) / temperature) next_token = int(mx.random.categorical(probs.reshape(1, -1))[0].item()) else: next_token = int(mx.argmax(logits[0, -1, :])) generated.append(next_token) if next_token == tokenizer.eos_token_id: break # Report theoretical savings saved_layers_total = max_new_tokens * (SKIP_TO - EXIT_LAYER - 1) return generated, 0.95, saved_layers_total # ─── API Endpoints ─── @app.get("/") def root(): return {"service": "Rushd-Agent API", "status": "live", "model": "Rushd-Geo + Early Exit"} @app.get("/health") def health(): return {"ok": True, "status": "live"} @app.post("/v1/chat/completions") async def chat(request: QueryRequest): t_start = time.time() # 1. Route expert = request.expert if request.expert != "auto" else route_query(request.query) # 2. Tokenize tokens = list(tokenizer.encode(request.query)) if len(tokens) > 4096: tokens = tokens[:4096] input_ids = mx.array([tokens]) # 3. Generate with early exit generated, nci, saved_layers = infer_early(input_ids, request.max_tokens, request.temperature) # 4. Decode output = tokenizer.decode(generated[len(tokens):]) elapsed = (time.time() - t_start) * 1000 mode = "early" if nci > 0.6 else "full" # OpenAI-compatible response return { "id": f"chatcmpl-{int(time.time())}", "object": "chat.completion", "created": int(time.time()), "model": "rushd-agent", "choices": [{ "index": 0, "message": { "role": "assistant", "content": output, }, "finish_reason": "stop", }], "usage": { "prompt_tokens": len(tokens), "completion_tokens": len(generated) - len(tokens), "total_tokens": len(generated), }, "x_rushd": { "expert": expert, "mode": mode, "nci": round(nci, 4), "saved_layers": saved_layers, "elapsed_ms": round(elapsed, 1), } } @app.get("/v1/models") def list_models(): return { "object": "list", "data": [{ "id": "rushd-agent", "object": "model", "created": int(time.time()), "owned_by": "rushd", }] } # ─── Main ─── if __name__ == "__main__": print("\n✅ Rushd-Agent API Ready!", flush=True) print(" Port: 8777", flush=True) print(" Endpoints:", flush=True) print(" GET /v1/models", flush=True) print(" POST /v1/chat/completions (OpenAI-compatible)", flush=True) print(f" Expert: auto-routed via Qwen2.5:7b", flush=True) print(f" Early Exit: L{EXIT_LAYER} → L{SKIP_TO} (64% savings)", flush=True) print(f" Starting server...\n", flush=True) uvicorn.run(app, host="0.0.0.0", port=8777)