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.
- package/README.md +154 -0
- package/bin/laya.mjs +362 -0
- package/daemon/anchors.json +111 -0
- package/daemon/requirements.txt +4 -0
- package/daemon/server.py +239 -0
- package/daemon/start.sh +67 -0
- package/extensions/laya-router.ts +440 -0
- package/package.json +43 -0
- package/skills/laya-tune/SKILL.md +47 -0
package/daemon/server.py
ADDED
|
@@ -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")
|
package/daemon/start.sh
ADDED
|
@@ -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
|