File size: 5,920 Bytes
f0b7fa0 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 | """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)
|