rushd-agent / code /rushd_api.py
BinSaqban's picture
Upload code/rushd_api.py with huggingface_hub
f0b7fa0 verified
Raw
History Blame Contribute Delete
5.92 kB
"""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)