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)