pi-laya-mlx 0.1.0

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
@@ -0,0 +1,239 @@
1
+ """Laya local decision daemon for pi auto-router.
2
+
3
+ Architecture:
4
+ Uses Laya's mmBERT Multilingual Encoder to produce dense sentence embeddings,
5
+ then computes cosine similarity against canonical anchor exemplars from anchors.json
6
+ (Semantic Exemplar KNN).
7
+
8
+ Endpoints:
9
+ GET /health -> {"ok": true, "ready": bool}
10
+ POST /predict -> {"answers": {...}, "latency_ms": float}
11
+ POST /reload-anchors -> {"ok": true, "reloaded_ms": float, "counts": {...}}
12
+ GET /logs -> {"logs": [...]}
13
+ POST /corrections -> {"ok": true, "total": int}
14
+ """
15
+
16
+ from datetime import datetime, timezone
17
+ import json
18
+ import os
19
+ import sys
20
+ import time
21
+ from typing import Dict, List
22
+
23
+ from fastapi import FastAPI, Request
24
+ from fastapi.responses import JSONResponse
25
+ import numpy as np
26
+
27
+ PORT = int(os.environ.get("LAYA_PORT", "4141"))
28
+ MODEL_ID = os.environ.get("LAYA_MODEL", "aac6fef/laya-multilingual-mlx")
29
+ BASE_DIR = os.path.dirname(__file__)
30
+ LOG_PATH = os.path.join(BASE_DIR, "predictions.jsonl")
31
+ CORRECTIONS_PATH = os.path.join(BASE_DIR, "corrections.jsonl")
32
+ ANCHORS_PATH = os.path.join(BASE_DIR, "anchors.json")
33
+
34
+ app = FastAPI(title="Laya MLX Local Daemon")
35
+ agent = None
36
+ embed_fn = None
37
+ load_error = None
38
+ anchor_mat: Dict[str, np.ndarray] = {}
39
+ current_anchors: Dict[str, List[str]] = {}
40
+
41
+
42
+ def _load_anchors_file() -> Dict[str, List[str]]:
43
+ if os.path.exists(ANCHORS_PATH):
44
+ try:
45
+ with open(ANCHORS_PATH, "r", encoding="utf-8") as f:
46
+ return json.load(f)
47
+ except Exception as e:
48
+ print(f"[laya-server] error reading anchors.json: {e}", file=sys.stderr)
49
+ return {}
50
+
51
+
52
+ def _init_anchors():
53
+ global anchor_mat, current_anchors
54
+ t0 = time.perf_counter()
55
+ anchors = _load_anchors_file()
56
+ if not anchors:
57
+ print("[laya-server] warning: anchors.json is empty or missing!", file=sys.stderr)
58
+ return 0.0
59
+
60
+ mat = {}
61
+ for cat, texts in anchors.items():
62
+ if texts:
63
+ vecs = embed_fn(texts) # (N, D)
64
+ norms = np.linalg.norm(vecs, axis=1, keepdims=True) + 1e-9
65
+ mat[cat] = vecs / norms
66
+ else:
67
+ mat[cat] = np.zeros((0, 768), dtype=np.float32)
68
+
69
+ anchor_mat = mat
70
+ current_anchors = anchors
71
+ ms = (time.perf_counter() - t0) * 1000
72
+ print(f"[laya-server] computed {sum(len(v) for v in anchors.values())} anchor embeddings across {len(anchors)} categories in {ms:.1f}ms", flush=True)
73
+ return ms
74
+
75
+
76
+ @app.on_event("startup")
77
+ def _load_model():
78
+ global agent, embed_fn, load_error
79
+ t0 = time.perf_counter()
80
+ try:
81
+ import laya_mlx as laya
82
+ agent = laya.load(MODEL_ID, dtype="float16")
83
+ embed_fn = laya.embed_fn_from_agent(agent)
84
+ _init_anchors()
85
+ # Warmup query
86
+ embed_fn(["warmup query"])
87
+ ms = (time.perf_counter() - t0) * 1000
88
+ print(f"[laya-server] model and encoder ready in {ms:.0f} ms: {MODEL_ID}", flush=True)
89
+ except Exception as e: # noqa: BLE001
90
+ agent = None
91
+ embed_fn = None
92
+ load_error = str(e)
93
+ print(f"[laya-server] model load FAILED: {e}", file=sys.stderr, flush=True)
94
+
95
+
96
+ @app.get("/health")
97
+ def health():
98
+ return {
99
+ "ok": agent is not None and embed_fn is not None,
100
+ "model": MODEL_ID,
101
+ "engine": "laya-multilingual-embeddings-knn",
102
+ "anchors_count": {k: len(v) for k, v in current_anchors.items()},
103
+ "error": load_error,
104
+ }
105
+
106
+
107
+ @app.post("/reload-anchors")
108
+ def reload_anchors():
109
+ """Hot-reload anchor embeddings from anchors.json in <100ms without restarting the server."""
110
+ if embed_fn is None:
111
+ return JSONResponse({"error": "model not loaded", "detail": load_error}, status_code=503)
112
+ ms = _init_anchors()
113
+ return {
114
+ "ok": True,
115
+ "reloaded_ms": round(ms, 2),
116
+ "counts": {k: len(v) for k, v in current_anchors.items()},
117
+ }
118
+
119
+
120
+ @app.get("/logs")
121
+ def get_logs(limit: int = 15):
122
+ """Return the most recent prediction logs."""
123
+ if not os.path.exists(LOG_PATH):
124
+ return {"logs": []}
125
+ try:
126
+ with open(LOG_PATH, "r", encoding="utf-8") as f:
127
+ lines = f.readlines()
128
+ recent = [json.loads(line.strip()) for line in lines[-limit:] if line.strip()]
129
+ return {"logs": recent}
130
+ except Exception as e: # noqa: BLE001
131
+ return JSONResponse({"error": str(e)}, status_code=500)
132
+
133
+
134
+ @app.post("/corrections")
135
+ async def record_correction(req: Request):
136
+ """Save a user correction/flag for offline tuning."""
137
+ try:
138
+ body = json.loads(await req.body())
139
+ except Exception as e:
140
+ return JSONResponse({"error": f"bad json: {e}"}, status_code=400)
141
+
142
+ prompt = body.get("prompt", "").strip()
143
+ target = body.get("target", "").strip()
144
+ if not prompt or not target:
145
+ return JSONResponse({"error": "prompt and target required"}, status_code=400)
146
+
147
+ entry = {
148
+ "timestamp": datetime.now(timezone.utc).isoformat(),
149
+ "prompt": prompt,
150
+ "target": target,
151
+ "was": body.get("was"),
152
+ }
153
+ try:
154
+ with open(CORRECTIONS_PATH, "a", encoding="utf-8") as f:
155
+ f.write(json.dumps(entry, ensure_ascii=False) + "\n")
156
+ except Exception as e:
157
+ return JSONResponse({"error": f"write failed: {e}"}, status_code=500)
158
+
159
+ return {"ok": True, "recorded": entry}
160
+
161
+
162
+ def _log_prediction(entry: dict):
163
+ try:
164
+ with open(LOG_PATH, "a", encoding="utf-8") as f:
165
+ f.write(json.dumps(entry, ensure_ascii=False) + "\n")
166
+ except Exception as e:
167
+ print(f"[laya-server] log error: {e}", file=sys.stderr)
168
+
169
+
170
+ def _classify_intent(prompt: str) -> dict:
171
+ pv = embed_fn([prompt]) # (1, D)
172
+ pv = pv / (np.linalg.norm(pv, axis=1, keepdims=True) + 1e-9)
173
+
174
+ cat_scores: Dict[str, float] = {}
175
+ for cat, mat in anchor_mat.items():
176
+ if mat.shape[0] == 0:
177
+ cat_scores[cat] = 0.0
178
+ continue
179
+ sims = np.dot(mat, pv.T).flatten() # (N,)
180
+ top3 = np.sort(sims)[-3:]
181
+ cat_scores[cat] = float(np.mean(top3))
182
+
183
+ # Softmax normalization with temperature scaling (T=0.05)
184
+ scores_arr = np.array(list(cat_scores.values()))
185
+ exp_scores = np.exp((scores_arr - np.max(scores_arr)) / 0.05)
186
+ probs_arr = exp_scores / np.sum(exp_scores)
187
+
188
+ prob_dict = {cat: round(float(p), 4) for cat, p in zip(cat_scores.keys(), probs_arr)}
189
+ top_cat = max(cat_scores, key=cat_scores.get)
190
+
191
+ return {
192
+ "type": "choice",
193
+ "choice": top_cat,
194
+ "confidence": prob_dict[top_cat],
195
+ "probabilities": prob_dict,
196
+ "raw_cosine": {k: round(v, 4) for k, v in cat_scores.items()},
197
+ }
198
+
199
+
200
+ @app.post("/predict")
201
+ async def predict(req: Request):
202
+ if agent is None or embed_fn is None:
203
+ return JSONResponse({"error": "model not loaded", "detail": load_error}, status_code=503)
204
+ try:
205
+ body = json.loads(await req.body())
206
+ except Exception as e: # noqa: BLE001
207
+ return JSONResponse({"error": f"bad json: {e}"}, status_code=400)
208
+
209
+ state = body.get("state", "")
210
+ if not state.strip():
211
+ return JSONResponse({"error": "state required"}, status_code=400)
212
+
213
+ t0 = time.perf_counter()
214
+ try:
215
+ classification = _classify_intent(state)
216
+ answers = {"task_category": classification}
217
+ except Exception as e: # noqa: BLE001
218
+ return JSONResponse({"error": f"predict failed: {e}"}, status_code=500)
219
+ ms = (time.perf_counter() - t0) * 1000
220
+
221
+ # Structured evaluation logging
222
+ entry = {
223
+ "timestamp": datetime.now(timezone.utc).isoformat(),
224
+ "prompt": state,
225
+ "choice": classification["choice"],
226
+ "confidence": classification["confidence"],
227
+ "probabilities": classification["probabilities"],
228
+ "raw_cosine": classification["raw_cosine"],
229
+ "latency_ms": round(ms, 2),
230
+ }
231
+ _log_prediction(entry)
232
+
233
+ return {"answers": answers, "latency_ms": round(ms, 2)}
234
+
235
+
236
+ if __name__ == "__main__":
237
+ import uvicorn
238
+
239
+ uvicorn.run(app, host="127.0.0.1", port=PORT, log_level="warning")
@@ -0,0 +1,67 @@
1
+ #!/bin/bash
2
+ # Start/stop/status/reload the Laya local decision daemon (port 4141).
3
+ # Usage: ./start.sh [start|stop|status|restart|reload]
4
+
5
+ DIR="$(cd "$(dirname "$0")" && pwd)"
6
+ PORT="${LAYA_PORT:-4141}"
7
+
8
+ is_running() {
9
+ curl -s -m 1 -o /dev/null "http://127.0.0.1:$PORT/health" && return 0 || return 1
10
+ }
11
+
12
+ case "${1:-start}" in
13
+ start)
14
+ if is_running; then
15
+ echo "Laya daemon already running on :$PORT"
16
+ curl -s "http://127.0.0.1:$PORT/health"; echo
17
+ exit 0
18
+ fi
19
+ cd "$DIR"
20
+ if command -v uv >/dev/null 2>&1; then
21
+ nohup uv run server.py > server.log 2>&1 &
22
+ else
23
+ nohup python3 server.py > server.log 2>&1 &
24
+ fi
25
+ echo "Laya daemon starting (pid $!)... waiting for model load"
26
+ for _ in $(seq 1 35); do
27
+ sleep 1
28
+ if is_running; then
29
+ echo "ready."
30
+ curl -s "http://127.0.0.1:$PORT/health"; echo
31
+ exit 0
32
+ fi
33
+ done
34
+ echo "failed to start — check $DIR/server.log"
35
+ exit 1
36
+ ;;
37
+ stop)
38
+ pkill -f "daemon/server.py" 2>/dev/null
39
+ sleep 1
40
+ lsof -ti:"$PORT" | xargs kill -9 2>/dev/null
41
+ echo "Laya daemon stopped."
42
+ ;;
43
+ status)
44
+ if is_running; then
45
+ echo "running:"
46
+ curl -s "http://127.0.0.1:$PORT/health"; echo
47
+ else
48
+ echo "not running."
49
+ exit 1
50
+ fi
51
+ ;;
52
+ reload)
53
+ if is_running; then
54
+ curl -s -X POST "http://127.0.0.1:$PORT/reload-anchors"; echo
55
+ else
56
+ echo "Laya daemon is not running."
57
+ exit 1
58
+ fi
59
+ ;;
60
+ restart)
61
+ "$0" stop; "$0" start
62
+ ;;
63
+ *)
64
+ echo "Usage: $0 [start|stop|status|restart|reload]"
65
+ exit 1
66
+ ;;
67
+ esac