sqljev 0.1.0__py3-none-any.whl
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.
- sqljev/__init__.py +6 -0
- sqljev/__main__.py +3 -0
- sqljev/aws_lambda.py +16 -0
- sqljev/cli.py +516 -0
- sqljev/core.py +674 -0
- sqljev/demo.py +307 -0
- sqljev/duckdb.py +56 -0
- sqljev/finetune.py +367 -0
- sqljev/gateway.py +215 -0
- sqljev/spark.py +64 -0
- sqljev-0.1.0.dist-info/METADATA +450 -0
- sqljev-0.1.0.dist-info/RECORD +16 -0
- sqljev-0.1.0.dist-info/WHEEL +4 -0
- sqljev-0.1.0.dist-info/entry_points.txt +2 -0
- sqljev-0.1.0.dist-info/licenses/LICENSE +176 -0
- sqljev-0.1.0.dist-info/licenses/NOTICE +36 -0
sqljev/finetune.py
ADDED
|
@@ -0,0 +1,367 @@
|
|
|
1
|
+
"""Fine-tune Laya on labelled SQL rows, on one GPU (a free Colab T4 is enough), and publish the checkpoint.
|
|
2
|
+
|
|
3
|
+
sqljev dataset "$DB_URL" "SELECT subject, body, team FROM tickets" --label team \\
|
|
4
|
+
--choice "which team should handle this?" --test-fraction 0.2 -o tickets.jsonl
|
|
5
|
+
sqljev finetune tickets.train.jsonl --out checkpoints/tickets --epochs 3
|
|
6
|
+
sqljev eval tickets.test.jsonl --model checkpoints/tickets
|
|
7
|
+
sqljev publish checkpoints/tickets --repo your-org/laya-tickets # Hugging Face Hub
|
|
8
|
+
|
|
9
|
+
Training data is what `sqljev dataset` writes: one row per line with the state (the row, as the model sees it at
|
|
10
|
+
query time), the question built by core.laya_question and the expected answer, so the checkpoint learns exactly
|
|
11
|
+
the questions sqljev will ask it.
|
|
12
|
+
|
|
13
|
+
The training loop follows Laya's own fine-tuning notebook (Convai Innovations, Apache 2.0,
|
|
14
|
+
https://github.com/NandhaKishorM/laya/tree/main/notebooks): policy-gradient updates against a proper scoring
|
|
15
|
+
rule plus soft cross-entropy, then per-question-type temperature calibration on a held-out slice. It is
|
|
16
|
+
single-GPU here, with gradient checkpointing and mixed precision, so it fits a 16 GB T4.
|
|
17
|
+
"""
|
|
18
|
+
import json
|
|
19
|
+
import os
|
|
20
|
+
import random
|
|
21
|
+
import re
|
|
22
|
+
import time
|
|
23
|
+
|
|
24
|
+
from .core import JevError, check_question, laya_question, to_row_json
|
|
25
|
+
|
|
26
|
+
LAYA_BASE = "convaiinnovations/laya"
|
|
27
|
+
SUBFOLDERS = {"english": None, "multilingual": "multilingual", "typed-decisions": "typed-decisions"}
|
|
28
|
+
|
|
29
|
+
|
|
30
|
+
# ---------------------------------------------------------------- examples
|
|
31
|
+
|
|
32
|
+
def examples_from_rows(rows, question, kind, options=None, label=None, columns=None, drop_nulls=True):
|
|
33
|
+
"""Labelled rows (dicts) -> training/eval records. The label column never reaches the state. For choice
|
|
34
|
+
questions without options, the options are the distinct labels."""
|
|
35
|
+
rows = [r for r in rows if r.get(label) is not None]
|
|
36
|
+
if kind == "choice" and not options:
|
|
37
|
+
options = sorted({str(r[label]) for r in rows})
|
|
38
|
+
kind, opts = check_question(kind, options)
|
|
39
|
+
q = {"q": laya_question(kind, question, opts)}
|
|
40
|
+
out = []
|
|
41
|
+
for r in rows:
|
|
42
|
+
lab = r[label]
|
|
43
|
+
if kind == "noul":
|
|
44
|
+
expected = _truthy(lab)
|
|
45
|
+
elif kind == "choice":
|
|
46
|
+
expected = str(lab)
|
|
47
|
+
if expected not in opts:
|
|
48
|
+
continue
|
|
49
|
+
else:
|
|
50
|
+
expected = opts.index(str(lab)) if str(lab) in opts else int(lab)
|
|
51
|
+
view = {k: v for k, v in r.items() if k != label and (not columns or k in columns)}
|
|
52
|
+
out.append({"state": json.loads(to_row_json(view, drop_nulls)), "questions": q, "expected": {"q": expected}})
|
|
53
|
+
return out
|
|
54
|
+
|
|
55
|
+
|
|
56
|
+
def split(records, test_fraction=0.2, seed=0):
|
|
57
|
+
"""Deterministic train/test split."""
|
|
58
|
+
idx = list(range(len(records)))
|
|
59
|
+
random.Random(seed).shuffle(idx)
|
|
60
|
+
n_test = int(len(records) * test_fraction)
|
|
61
|
+
test = set(idx[:n_test])
|
|
62
|
+
return [r for i, r in enumerate(records) if i not in test], [r for i, r in enumerate(records) if i in test]
|
|
63
|
+
|
|
64
|
+
|
|
65
|
+
def write_jsonl(records, path):
|
|
66
|
+
with open(path, "w") as f:
|
|
67
|
+
for r in records:
|
|
68
|
+
f.write(json.dumps(r, ensure_ascii=False, default=str) + "\n")
|
|
69
|
+
|
|
70
|
+
|
|
71
|
+
def read_jsonl(path):
|
|
72
|
+
with open(path) as f:
|
|
73
|
+
return [json.loads(line) for line in f if line.strip() and not line.startswith("#")]
|
|
74
|
+
|
|
75
|
+
|
|
76
|
+
def _truthy(v):
|
|
77
|
+
if isinstance(v, bool):
|
|
78
|
+
return v
|
|
79
|
+
s = str(v).strip().lower()
|
|
80
|
+
if s in ("1", "true", "t", "yes", "y"):
|
|
81
|
+
return True
|
|
82
|
+
if s in ("0", "false", "f", "no", "n"):
|
|
83
|
+
return False
|
|
84
|
+
raise JevError("sqljev: label %r is not a boolean" % (v,))
|
|
85
|
+
|
|
86
|
+
|
|
87
|
+
def accuracy(records, **settings):
|
|
88
|
+
"""Share of expected answers the model gets right on `records` (a list or a JSONL path), with any engine
|
|
89
|
+
settings (model=<checkpoint>, device=..., backend=...). Returns {"accuracy", "decisions", "seconds"}."""
|
|
90
|
+
from .core import Jev
|
|
91
|
+
records = read_jsonl(records) if isinstance(records, str) else records
|
|
92
|
+
jev = Jev(**settings)
|
|
93
|
+
groups = {}
|
|
94
|
+
for rec in records:
|
|
95
|
+
for qid, q in rec["questions"].items():
|
|
96
|
+
crit = q.get("criteria")
|
|
97
|
+
opts = list(crit) if crit else None
|
|
98
|
+
groups.setdefault((q["type"], q["instructions"], json.dumps(opts)), []).append(
|
|
99
|
+
(rec["state"], rec["expected"][qid]))
|
|
100
|
+
t0, right, total = time.time(), 0, 0
|
|
101
|
+
for (kind, instr, opts_json), items in groups.items():
|
|
102
|
+
prefix = "Is it true that "
|
|
103
|
+
query = instr[len(prefix):-1] if kind == "noul" and instr.startswith(prefix) and instr.endswith("?") else instr
|
|
104
|
+
answers = jev.evaluate([s for s, _ in items], query, kind, json.loads(opts_json))
|
|
105
|
+
for (_, exp), a in zip(items, answers):
|
|
106
|
+
if kind == "noul":
|
|
107
|
+
right += (a["noul"] >= 0.5) == bool(exp)
|
|
108
|
+
elif kind == "choice":
|
|
109
|
+
right += a["choice"] == exp
|
|
110
|
+
else:
|
|
111
|
+
right += round(a["score"]) == int(exp)
|
|
112
|
+
total += 1
|
|
113
|
+
return {"accuracy": round(right / total, 4) if total else 0.0, "decisions": total,
|
|
114
|
+
"seconds": round(time.time() - t0, 1)}
|
|
115
|
+
|
|
116
|
+
|
|
117
|
+
# ---------------------------------------------------------------- training
|
|
118
|
+
|
|
119
|
+
def base_dir(base=LAYA_BASE):
|
|
120
|
+
"""Local directory of a Laya checkpoint: a path, a named base checkpoint, or a Hub id."""
|
|
121
|
+
if os.path.isdir(base):
|
|
122
|
+
return base
|
|
123
|
+
from huggingface_hub import snapshot_download
|
|
124
|
+
sub = SUBFOLDERS.get(base, None)
|
|
125
|
+
repo = LAYA_BASE if base in SUBFOLDERS else base
|
|
126
|
+
d = snapshot_download(repo)
|
|
127
|
+
try:
|
|
128
|
+
from laya.agent import _fix_tokenizer_config
|
|
129
|
+
_fix_tokenizer_config(d)
|
|
130
|
+
except Exception: # noqa: BLE001 -- older/newer laya: the files are already usable
|
|
131
|
+
pass
|
|
132
|
+
return os.path.join(d, sub) if sub else d
|
|
133
|
+
|
|
134
|
+
|
|
135
|
+
def _items(records, tok, cfg):
|
|
136
|
+
from laya.common import QTYPES, build_sequence, render_options
|
|
137
|
+
items, skipped = [], 0
|
|
138
|
+
for rec in records:
|
|
139
|
+
for qid, q in rec["questions"].items():
|
|
140
|
+
t, crit = q["type"], q.get("criteria") or {}
|
|
141
|
+
exp = rec["expected"][qid]
|
|
142
|
+
if t == "noul":
|
|
143
|
+
target = [0.0, 1.0] if exp else [1.0, 0.0]
|
|
144
|
+
elif t == "choice":
|
|
145
|
+
keys = list(crit.keys()) if isinstance(crit, dict) else list(crit)
|
|
146
|
+
target = [float(k == exp) for k in keys]
|
|
147
|
+
else:
|
|
148
|
+
target = [float(i == int(exp)) for i in range(len(crit))]
|
|
149
|
+
seq, markers = build_sequence(tok, rec["state"], {"t": t, "ins": q["instructions"], "crit": crit},
|
|
150
|
+
cfg["max_len"], cfg["head_max_len"])
|
|
151
|
+
if len(markers) != len(render_options({"t": t, "crit": crit})) or sum(target) != 1:
|
|
152
|
+
skipped += 1
|
|
153
|
+
continue
|
|
154
|
+
items.append({"ids": seq, "markers": markers, "qtype": QTYPES[t], "target": target,
|
|
155
|
+
"label": target.index(1.0)})
|
|
156
|
+
return items, skipped
|
|
157
|
+
|
|
158
|
+
|
|
159
|
+
def _collate(items, pad_id):
|
|
160
|
+
import torch
|
|
161
|
+
n, L = len(items), max(len(it["ids"]) for it in items)
|
|
162
|
+
kmax = max(len(it["markers"]) for it in items)
|
|
163
|
+
ids = torch.full((n, L), pad_id, dtype=torch.long)
|
|
164
|
+
att = torch.zeros((n, L), dtype=torch.long)
|
|
165
|
+
mpos = torch.zeros((n, kmax), dtype=torch.long)
|
|
166
|
+
mmask = torch.zeros((n, kmax), dtype=torch.bool)
|
|
167
|
+
target = torch.zeros((n, kmax), dtype=torch.float32)
|
|
168
|
+
for i, it in enumerate(items):
|
|
169
|
+
ids[i, :len(it["ids"])] = torch.tensor(it["ids"])
|
|
170
|
+
att[i, :len(it["ids"])] = 1
|
|
171
|
+
k = len(it["markers"])
|
|
172
|
+
mpos[i, :k] = torch.tensor(it["markers"])
|
|
173
|
+
mmask[i, :k] = True
|
|
174
|
+
target[i, :k] = torch.tensor(it["target"])
|
|
175
|
+
return ids, att, mpos, mmask, target, torch.tensor([it["qtype"] for it in items])
|
|
176
|
+
|
|
177
|
+
|
|
178
|
+
def _fit_temperature(pairs):
|
|
179
|
+
import torch
|
|
180
|
+
if len(pairs) < 10:
|
|
181
|
+
return 1.0
|
|
182
|
+
kmax = max(len(z) for z, _ in pairs)
|
|
183
|
+
Z = torch.full((len(pairs), kmax), -1e4)
|
|
184
|
+
T = torch.zeros((len(pairs), kmax))
|
|
185
|
+
for i, (z, t) in enumerate(pairs):
|
|
186
|
+
Z[i, :len(z)] = torch.tensor(z)
|
|
187
|
+
T[i, :len(t)] = torch.tensor(t)
|
|
188
|
+
log_t = torch.zeros(1, requires_grad=True)
|
|
189
|
+
opt = torch.optim.LBFGS([log_t], lr=0.1, max_iter=100)
|
|
190
|
+
|
|
191
|
+
def closure():
|
|
192
|
+
opt.zero_grad()
|
|
193
|
+
loss = -(T * torch.log_softmax(Z / log_t.exp(), -1)).sum(-1).mean()
|
|
194
|
+
loss.backward()
|
|
195
|
+
return loss
|
|
196
|
+
opt.step(closure)
|
|
197
|
+
return float(torch.clamp(log_t.exp(), 0.5, 5.0).item())
|
|
198
|
+
|
|
199
|
+
|
|
200
|
+
def finetune(train, out_dir, base=LAYA_BASE, epochs=3, micro_batch=8, grad_accum=8, lr_encoder=2.5e-5,
|
|
201
|
+
lr_head=1e-4, group_size=4, train_layers=None, device=None, seed=0, log=print):
|
|
202
|
+
"""Fine-tune a Laya checkpoint on `train` (records or a JSONL path) and save it to `out_dir` in Laya's
|
|
203
|
+
format, loadable with laya.load(out_dir) and usable as SQLJEV_MODEL=out_dir.
|
|
204
|
+
|
|
205
|
+
train_layers: train only the top N encoder layers plus the decision head, freezing the embeddings and the
|
|
206
|
+
layers below (low-memory mode for GPUs under ~10 GB free). None trains everything."""
|
|
207
|
+
import torch
|
|
208
|
+
from laya.common import build_model, proper_reward
|
|
209
|
+
from safetensors.torch import load_file, save_file
|
|
210
|
+
from transformers import AutoTokenizer
|
|
211
|
+
|
|
212
|
+
records = read_jsonl(train) if isinstance(train, str) else list(train)
|
|
213
|
+
if not records:
|
|
214
|
+
raise JevError("sqljev: no training records")
|
|
215
|
+
device = torch.device(device or ("cuda" if torch.cuda.is_available() else "cpu"))
|
|
216
|
+
src = base_dir(base)
|
|
217
|
+
with open(os.path.join(src, "rl_agent_config.json")) as f:
|
|
218
|
+
cfg = json.load(f)
|
|
219
|
+
tok = AutoTokenizer.from_pretrained(os.path.join(src, "tokenizer"))
|
|
220
|
+
items, skipped = _items(records, tok, cfg)
|
|
221
|
+
if not items:
|
|
222
|
+
raise JevError("sqljev: none of the records could be encoded for this checkpoint")
|
|
223
|
+
rng = random.Random(seed)
|
|
224
|
+
rng.shuffle(items)
|
|
225
|
+
n_cal = min(400, len(items) // 10)
|
|
226
|
+
calib, items = items[:n_cal], items[n_cal:]
|
|
227
|
+
log("sqljev finetune: %d training items, %d held out for calibration%s, base %s, device %s"
|
|
228
|
+
% (len(items), len(calib), ", %d skipped" % skipped if skipped else "", base, device))
|
|
229
|
+
|
|
230
|
+
torch.manual_seed(seed)
|
|
231
|
+
model = build_model(cfg, encoder_dir=os.path.join(src, "encoder"))
|
|
232
|
+
model.load_state_dict(load_file(os.path.join(src, "model.safetensors")), strict=True)
|
|
233
|
+
if device.type == "cuda":
|
|
234
|
+
model.encoder.gradient_checkpointing_enable(gradient_checkpointing_kwargs={"use_reentrant": False})
|
|
235
|
+
model.head_checkpointing = True
|
|
236
|
+
model.to(device).train()
|
|
237
|
+
if train_layers is not None:
|
|
238
|
+
n_layers = 1 + max(int(m.group(1)) for n, _ in model.named_parameters()
|
|
239
|
+
for m in [re.match(r"encoder\.layers\.(\d+)\.", n)] if m)
|
|
240
|
+
keep = n_layers - int(train_layers)
|
|
241
|
+
for n, p in model.named_parameters():
|
|
242
|
+
m = re.match(r"encoder\.layers\.(\d+)\.", n)
|
|
243
|
+
if n.startswith("encoder.embeddings.") or (m and int(m.group(1)) < keep):
|
|
244
|
+
p.requires_grad_(False)
|
|
245
|
+
log("sqljev finetune: training the top %d of %d encoder layers and the head" % (min(int(train_layers), n_layers), n_layers))
|
|
246
|
+
|
|
247
|
+
use_cuda = device.type == "cuda"
|
|
248
|
+
amp_dtype = torch.bfloat16 if use_cuda and torch.cuda.is_bf16_supported() else torch.float16
|
|
249
|
+
scaler = torch.amp.GradScaler("cuda", enabled=use_cuda and amp_dtype == torch.float16)
|
|
250
|
+
enc = [p for n, p in model.named_parameters() if n.startswith("encoder.") and p.requires_grad]
|
|
251
|
+
head = [p for n, p in model.named_parameters() if not n.startswith("encoder.") and p.requires_grad]
|
|
252
|
+
opt = torch.optim.AdamW([{"params": enc, "lr": lr_encoder}, {"params": head, "lr": lr_head}], weight_decay=0.01)
|
|
253
|
+
updates = max(1, -(-len(items) // (micro_batch * grad_accum)) * epochs)
|
|
254
|
+
sched = torch.optim.lr_scheduler.CosineAnnealingLR(opt, T_max=updates, eta_min=1e-6)
|
|
255
|
+
t0 = time.time()
|
|
256
|
+
for epoch in range(epochs):
|
|
257
|
+
rng.shuffle(items)
|
|
258
|
+
sigma = 0.4 + (0.1 - 0.4) * (epoch / max(1, epochs - 1))
|
|
259
|
+
total, steps = 0.0, 0
|
|
260
|
+
opt.zero_grad(set_to_none=True)
|
|
261
|
+
for b in range(0, len(items), micro_batch):
|
|
262
|
+
ids, att, mpos, mmask, target, qtype = (x.to(device) for x in _collate(items[b:b + micro_batch], tok.pad_token_id))
|
|
263
|
+
with torch.autocast(device.type, dtype=amp_dtype, enabled=use_cuda):
|
|
264
|
+
logits, act = model(ids, att, mpos, mmask, qtype)
|
|
265
|
+
logits = logits.float()
|
|
266
|
+
k = mmask.sum(-1, keepdim=True).float()
|
|
267
|
+
eps = torch.randn((group_size,) + logits.shape, device=device) * sigma * mmask
|
|
268
|
+
eps = (eps - eps.sum(-1, keepdim=True) / k) * mmask
|
|
269
|
+
z = logits.detach().unsqueeze(0) + eps
|
|
270
|
+
q = torch.softmax(z.masked_fill(~mmask, -1e4), -1)
|
|
271
|
+
with torch.no_grad():
|
|
272
|
+
r = proper_reward(q, target.unsqueeze(0), qtype, mmask, w_sph=0.75, w_rps=1.0)
|
|
273
|
+
adv = (r - r.mean(0, keepdim=True))
|
|
274
|
+
adv = adv / (adv.std() + 1e-6)
|
|
275
|
+
logp = -(((z - logits.unsqueeze(0)) ** 2) * mmask).sum(-1) / (2 * sigma ** 2)
|
|
276
|
+
loss_ce = -(target * torch.log_softmax(logits.masked_fill(~mmask, -1e4), -1)).sum(-1).mean()
|
|
277
|
+
loss = (-(adv * logp).mean() + loss_ce) / grad_accum + 0.0 * act.sum()
|
|
278
|
+
scaler.scale(loss).backward()
|
|
279
|
+
steps += 1
|
|
280
|
+
total += loss.item() * grad_accum
|
|
281
|
+
if steps % grad_accum == 0 or b + micro_batch >= len(items):
|
|
282
|
+
scaler.unscale_(opt)
|
|
283
|
+
torch.nn.utils.clip_grad_norm_(enc + head, 1.0)
|
|
284
|
+
scaler.step(opt)
|
|
285
|
+
scaler.update()
|
|
286
|
+
sched.step()
|
|
287
|
+
opt.zero_grad(set_to_none=True)
|
|
288
|
+
if steps % 100 == 0:
|
|
289
|
+
log(" epoch %d/%d step %d/%d loss %.4f %.0fs"
|
|
290
|
+
% (epoch + 1, epochs, steps, -(-len(items) // micro_batch), total / steps, time.time() - t0))
|
|
291
|
+
log(" epoch %d/%d done, mean loss %.4f, %.0fs" % (epoch + 1, epochs, total / max(1, steps), time.time() - t0))
|
|
292
|
+
|
|
293
|
+
# Calibrate one temperature per question type on items the run never trained on.
|
|
294
|
+
model.eval()
|
|
295
|
+
preds = []
|
|
296
|
+
with torch.no_grad():
|
|
297
|
+
for b in range(0, len(calib), 16):
|
|
298
|
+
chunk = calib[b:b + 16]
|
|
299
|
+
ids, att, mpos, mmask, target, qtype = (x.to(device) for x in _collate(chunk, tok.pad_token_id))
|
|
300
|
+
with torch.autocast(device.type, dtype=amp_dtype, enabled=use_cuda):
|
|
301
|
+
lg, _ = model(ids, att, mpos, mmask, qtype)
|
|
302
|
+
lg = lg.float().cpu().numpy()
|
|
303
|
+
preds += [(it["qtype"], lg[i, :len(it["markers"])].tolist(), it["target"]) for i, it in enumerate(chunk)]
|
|
304
|
+
temps = list(cfg.get("temperature") or [1.2, 1.2, 1.2])
|
|
305
|
+
for qt in range(3):
|
|
306
|
+
sel = [(z, t) for q, z, t in preds if q == qt]
|
|
307
|
+
if len(sel) >= 10:
|
|
308
|
+
temps[qt] = _fit_temperature(sel)
|
|
309
|
+
|
|
310
|
+
os.makedirs(out_dir, exist_ok=True)
|
|
311
|
+
save_file({k: v.half().contiguous().cpu() for k, v in model.state_dict().items()},
|
|
312
|
+
os.path.join(out_dir, "model.safetensors"))
|
|
313
|
+
model.encoder.config.save_pretrained(os.path.join(out_dir, "encoder"))
|
|
314
|
+
tok.save_pretrained(os.path.join(out_dir, "tokenizer"))
|
|
315
|
+
cfg.update(fine_tuned=True, model_name="laya-sqljev", temperature=temps)
|
|
316
|
+
cfg.pop("temperature_by_options", None)
|
|
317
|
+
with open(os.path.join(out_dir, "rl_agent_config.json"), "w") as f:
|
|
318
|
+
json.dump(cfg, f, indent=2)
|
|
319
|
+
questions = sorted({q["instructions"] for r in records for q in r["questions"].values()})
|
|
320
|
+
with open(os.path.join(out_dir, "sqljev_finetune.json"), "w") as f:
|
|
321
|
+
json.dump({"base": base, "records": len(records), "items": len(items), "epochs": epochs,
|
|
322
|
+
"questions": questions, "seconds": round(time.time() - t0, 1)}, f, indent=2)
|
|
323
|
+
log("sqljev finetune: saved to %s (%.0fs)" % (out_dir, time.time() - t0))
|
|
324
|
+
return out_dir
|
|
325
|
+
|
|
326
|
+
|
|
327
|
+
# ---------------------------------------------------------------- publishing
|
|
328
|
+
|
|
329
|
+
def publish(out_dir, repo_id, token=None, private=True, metrics=None):
|
|
330
|
+
"""Upload a fine-tuned checkpoint to the Hugging Face Hub with a model card. Every database then uses it
|
|
331
|
+
with SQLJEV_MODEL=<repo_id>."""
|
|
332
|
+
from huggingface_hub import HfApi
|
|
333
|
+
meta_path = os.path.join(out_dir, "sqljev_finetune.json")
|
|
334
|
+
meta = json.load(open(meta_path)) if os.path.exists(meta_path) else {}
|
|
335
|
+
if metrics:
|
|
336
|
+
meta["metrics"] = metrics
|
|
337
|
+
with open(meta_path, "w") as f:
|
|
338
|
+
json.dump(meta, f, indent=2)
|
|
339
|
+
card = """---
|
|
340
|
+
license: apache-2.0
|
|
341
|
+
base_model: %s
|
|
342
|
+
tags: [laya, sqljev, sql, decision-model]
|
|
343
|
+
---
|
|
344
|
+
|
|
345
|
+
# %s
|
|
346
|
+
|
|
347
|
+
A [Laya](https://github.com/NandhaKishorM/laya) checkpoint fine-tuned with [sqljev](https://github.com/singhpratech/sqljev)
|
|
348
|
+
on labelled SQL rows.
|
|
349
|
+
|
|
350
|
+
Questions it was trained on:
|
|
351
|
+
%s
|
|
352
|
+
|
|
353
|
+
%s
|
|
354
|
+
Use it from any database sqljev supports:
|
|
355
|
+
|
|
356
|
+
```bash
|
|
357
|
+
SQLJEV_MODEL=%s sqljev gateway --host 0.0.0.0
|
|
358
|
+
```
|
|
359
|
+
""" % (meta.get("base", LAYA_BASE), repo_id.split("/")[-1],
|
|
360
|
+
"\n".join("- %s" % q for q in meta.get("questions", [])),
|
|
361
|
+
("Held-out accuracy: %s\n" % json.dumps(metrics)) if metrics else "", repo_id)
|
|
362
|
+
with open(os.path.join(out_dir, "README.md"), "w") as f:
|
|
363
|
+
f.write(card)
|
|
364
|
+
api = HfApi(token=token)
|
|
365
|
+
api.create_repo(repo_id, private=private, exist_ok=True)
|
|
366
|
+
api.upload_folder(repo_id=repo_id, folder_path=out_dir, commit_message="sqljev finetune")
|
|
367
|
+
return "https://huggingface.co/" + repo_id
|
sqljev/gateway.py
ADDED
|
@@ -0,0 +1,215 @@
|
|
|
1
|
+
"""sqljev gateway: one HTTP service that answers the batch protocols databases use to call out.
|
|
2
|
+
|
|
3
|
+
sqljev gateway --port 8765 # Laya in-process (backend 'local') by default
|
|
4
|
+
|
|
5
|
+
Routes (all POST, JSON):
|
|
6
|
+
|
|
7
|
+
/v1/eval sqljev native: {"question", "kind", "options", "rows": [...]} -> {"answers": [...]}
|
|
8
|
+
/sqlserver same body; called by jev.judge through sp_invoke_external_rest_endpoint
|
|
9
|
+
/bigquery BigQuery remote functions: {"calls": [[args]], "userDefinedContext": {"fn": ...}}
|
|
10
|
+
-> {"replies": [...]}
|
|
11
|
+
/snowflake/<fn> Snowflake external / service functions: {"data": [[rownum, args...]]}
|
|
12
|
+
-> {"data": [[rownum, result]]}
|
|
13
|
+
/redshift Redshift Lambda UDF payload over HTTP (the same handler as sqljev.aws_lambda)
|
|
14
|
+
/v1/systemone TypeSafe Jev's wire protocol, for Jev clients such as pg-jev: batch requests
|
|
15
|
+
{"state": {"condition", "rows": [...]}, "questions": {"r0": ...}} are answered by the
|
|
16
|
+
engine, so `SET jev.api_url = 'http://gateway:8765/v1/systemone'` runs pg-jev on Laya
|
|
17
|
+
GET /health, GET /stats
|
|
18
|
+
|
|
19
|
+
Each database already sends rows in batches (BigQuery up to max_batching_rows, Snowflake/Redshift in their
|
|
20
|
+
own batch sizes), so every request becomes one engine.call(): rows grouped by question, de-duplicated,
|
|
21
|
+
looked up in the cache and the misses judged in shared forward passes.
|
|
22
|
+
|
|
23
|
+
Auth: set SQLJEV_GATEWAY_TOKEN and send it as "Authorization: Bearer <token>" or "X-Jev-Token: <token>".
|
|
24
|
+
BigQuery cannot send headers: run the gateway on Cloud Run with IAM auth instead.
|
|
25
|
+
"""
|
|
26
|
+
import hmac
|
|
27
|
+
import json
|
|
28
|
+
import os
|
|
29
|
+
import re
|
|
30
|
+
import sys
|
|
31
|
+
import traceback
|
|
32
|
+
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
|
|
33
|
+
|
|
34
|
+
from .core import FUNCTIONS, JevError, Jev, __version__, function_name
|
|
35
|
+
|
|
36
|
+
MAX_BODY_BYTES = 64 * 1024 * 1024
|
|
37
|
+
|
|
38
|
+
|
|
39
|
+
def handle_eval(engine, body):
|
|
40
|
+
kind = body.get("kind") or "noul"
|
|
41
|
+
rows = body.get("rows")
|
|
42
|
+
if not isinstance(rows, list) or not body.get("question"):
|
|
43
|
+
raise JevError("sqljev: body needs 'question' and a 'rows' array")
|
|
44
|
+
return {"answers": engine.evaluate(rows, body["question"], kind, body.get("options"))}
|
|
45
|
+
|
|
46
|
+
|
|
47
|
+
def handle_bigquery(engine, body):
|
|
48
|
+
ctx = body.get("userDefinedContext") or {}
|
|
49
|
+
fn = ctx.get("fn") or ctx.get("function")
|
|
50
|
+
if not fn:
|
|
51
|
+
raise JevError("sqljev: set user_defined_context = [('fn', 'jev_prob')] on the remote function")
|
|
52
|
+
return {"replies": _json_safe(engine.call(fn, body.get("calls") or []))}
|
|
53
|
+
|
|
54
|
+
|
|
55
|
+
def handle_snowflake(engine, body, fn):
|
|
56
|
+
data = body.get("data") or []
|
|
57
|
+
results = engine.call(fn, [r[1:] for r in data]) # VARIANT takes jev_eval's object as is
|
|
58
|
+
return {"data": [[r[0], v] for r, v in zip(data, results)]}
|
|
59
|
+
|
|
60
|
+
|
|
61
|
+
def handle_redshift(engine, event):
|
|
62
|
+
"""Redshift Lambda UDF protocol. Errors are reported in-band, as Redshift expects."""
|
|
63
|
+
try:
|
|
64
|
+
fn = event.get("external_function") or ""
|
|
65
|
+
args = event.get("arguments") or []
|
|
66
|
+
results = engine.call(fn, args)
|
|
67
|
+
return {"success": True, "num_records": len(results), "results": _json_safe(results)}
|
|
68
|
+
except Exception as e: # noqa: BLE001 -- Redshift shows error_msg to the user
|
|
69
|
+
return {"success": False, "error_msg": str(e)}
|
|
70
|
+
|
|
71
|
+
|
|
72
|
+
_ROW_REF = re.compile(r"`rows\[(\d+)\]`")
|
|
73
|
+
_JEV_PREFIX = {"score": "Rate the record `rows[%d]`: ", "choice": "For the record `rows[%d]`: "}
|
|
74
|
+
|
|
75
|
+
|
|
76
|
+
def handle_systemone(engine, body):
|
|
77
|
+
"""Answer a Jev batch request (one shared state with `rows`, one question per row, the format pg-jev and
|
|
78
|
+
sqljev's own jev backend send) with the engine: questions are grouped back into (question, rows) pairs."""
|
|
79
|
+
state, questions = body.get("state"), body.get("questions")
|
|
80
|
+
if not isinstance(state, dict) or not isinstance(state.get("rows"), list) or not isinstance(questions, dict):
|
|
81
|
+
raise JevError("sqljev: /v1/systemone accepts batch requests: state {'rows': [...]} with one question per "
|
|
82
|
+
"row that names `rows[i]` (as pg-jev sends); use /v1/eval for everything else")
|
|
83
|
+
rows, groups = state["rows"], {}
|
|
84
|
+
for qid, q in questions.items():
|
|
85
|
+
m = _ROW_REF.search(q.get("instructions") or "")
|
|
86
|
+
if not m or int(m.group(1)) >= len(rows):
|
|
87
|
+
raise JevError("sqljev: question %r does not name a row as `rows[i]`" % qid)
|
|
88
|
+
i, kind = int(m.group(1)), q.get("type")
|
|
89
|
+
if kind == "noul":
|
|
90
|
+
query, opts = state.get("condition") or "", None
|
|
91
|
+
elif kind in _JEV_PREFIX:
|
|
92
|
+
prefix = _JEV_PREFIX[kind] % i
|
|
93
|
+
instr = q["instructions"]
|
|
94
|
+
query = instr[len(prefix):] if instr.startswith(prefix) else instr
|
|
95
|
+
crit = q.get("criteria")
|
|
96
|
+
opts = list(crit.keys()) if isinstance(crit, dict) else list(crit or [])
|
|
97
|
+
else:
|
|
98
|
+
raise JevError("sqljev: unknown question type %r" % kind)
|
|
99
|
+
groups.setdefault((kind, query, tuple(opts or ())), []).append((qid, i))
|
|
100
|
+
answers = {}
|
|
101
|
+
for (kind, query, opts), items in groups.items():
|
|
102
|
+
got = engine.evaluate([rows[i] for _, i in items], query, kind, list(opts) or None)
|
|
103
|
+
answers.update({qid: a for (qid, _), a in zip(items, got)})
|
|
104
|
+
return {"model": "sqljev-" + engine.cfg["backend"], "answers": answers,
|
|
105
|
+
"usage": {"input_tokens": 0, "output_tokens": 0}}
|
|
106
|
+
|
|
107
|
+
|
|
108
|
+
def _json_safe(values):
|
|
109
|
+
# jev_eval answers go back to warehouses as JSON text columns (BigQuery JSON/STRING, Snowflake VARIANT
|
|
110
|
+
# accepts objects directly but a string is safe everywhere).
|
|
111
|
+
return [json.dumps(v) if isinstance(v, dict) else v for v in values]
|
|
112
|
+
|
|
113
|
+
|
|
114
|
+
class Handler(BaseHTTPRequestHandler):
|
|
115
|
+
protocol_version = "HTTP/1.1"
|
|
116
|
+
engine = None
|
|
117
|
+
token = None
|
|
118
|
+
server_version = "sqljev/" + __version__
|
|
119
|
+
|
|
120
|
+
def log_message(self, fmt, *args):
|
|
121
|
+
if os.environ.get("SQLJEV_GATEWAY_LOG"):
|
|
122
|
+
sys.stderr.write("sqljev gateway: " + fmt % args + "\n")
|
|
123
|
+
|
|
124
|
+
def _send(self, code, obj):
|
|
125
|
+
data = json.dumps(obj).encode()
|
|
126
|
+
self.send_response(code)
|
|
127
|
+
self.send_header("Content-Type", "application/json")
|
|
128
|
+
self.send_header("Content-Length", str(len(data)))
|
|
129
|
+
self.end_headers()
|
|
130
|
+
self.wfile.write(data)
|
|
131
|
+
|
|
132
|
+
def _authorized(self):
|
|
133
|
+
if not self.token:
|
|
134
|
+
return True
|
|
135
|
+
got = self.headers.get("X-Jev-Token") or ""
|
|
136
|
+
auth = self.headers.get("Authorization") or ""
|
|
137
|
+
if auth.startswith("Bearer "):
|
|
138
|
+
got = got or auth[7:]
|
|
139
|
+
return hmac.compare_digest(got.encode(), self.token.encode())
|
|
140
|
+
|
|
141
|
+
def do_GET(self):
|
|
142
|
+
if self.path == "/health":
|
|
143
|
+
return self._send(200, {"status": "ok", "version": __version__, "backend": self.engine.cfg["backend"],
|
|
144
|
+
"functions": sorted(FUNCTIONS)})
|
|
145
|
+
if self.path == "/stats":
|
|
146
|
+
if not self._authorized():
|
|
147
|
+
return self._send(401, {"error": "unauthorized"})
|
|
148
|
+
return self._send(200, self.engine.stats())
|
|
149
|
+
self._send(404, {"error": "not found"})
|
|
150
|
+
|
|
151
|
+
def do_POST(self):
|
|
152
|
+
n = int(self.headers.get("Content-Length") or 0)
|
|
153
|
+
if n > MAX_BODY_BYTES:
|
|
154
|
+
return self._send(413, {"error": "body too large"})
|
|
155
|
+
raw = self.rfile.read(n)
|
|
156
|
+
if not self._authorized():
|
|
157
|
+
return self._send(401, {"error": "unauthorized"})
|
|
158
|
+
path = self.path.split("?", 1)[0].rstrip("/")
|
|
159
|
+
try:
|
|
160
|
+
body = json.loads(raw or b"{}")
|
|
161
|
+
if path in ("/v1/eval", "/sqlserver"):
|
|
162
|
+
return self._send(200, handle_eval(self.engine, body))
|
|
163
|
+
if path == "/bigquery":
|
|
164
|
+
return self._send(200, handle_bigquery(self.engine, body))
|
|
165
|
+
if path.startswith("/snowflake/"):
|
|
166
|
+
return self._send(200, handle_snowflake(self.engine, body, function_name(path.rsplit("/", 1)[1])))
|
|
167
|
+
if path == "/redshift":
|
|
168
|
+
return self._send(200, handle_redshift(self.engine, body))
|
|
169
|
+
if path == "/v1/systemone":
|
|
170
|
+
return self._send(200, handle_systemone(self.engine, body))
|
|
171
|
+
return self._send(404, {"error": "not found: " + path})
|
|
172
|
+
except (JevError, ValueError, KeyError, TypeError, IndexError) as e:
|
|
173
|
+
# 400 is not retried by BigQuery/Snowflake: validation errors surface as query errors.
|
|
174
|
+
return self._send(400, {"error": str(e), "errorMessage": str(e)})
|
|
175
|
+
except Exception as e: # noqa: BLE001
|
|
176
|
+
traceback.print_exc()
|
|
177
|
+
return self._send(500, {"error": "internal error: %s" % type(e).__name__,
|
|
178
|
+
"errorMessage": "sqljev gateway internal error"})
|
|
179
|
+
|
|
180
|
+
|
|
181
|
+
def make_server(host="127.0.0.1", port=8765, engine=None, token=None):
|
|
182
|
+
handler = type("SqlJevHandler", (Handler,), {
|
|
183
|
+
"engine": engine or Jev(),
|
|
184
|
+
"token": token if token is not None else os.environ.get("SQLJEV_GATEWAY_TOKEN"),
|
|
185
|
+
})
|
|
186
|
+
ThreadingHTTPServer.daemon_threads = True
|
|
187
|
+
return ThreadingHTTPServer((host, port), handler)
|
|
188
|
+
|
|
189
|
+
|
|
190
|
+
def serve(host="127.0.0.1", port=8765, engine=None, token=None, certfile=None, keyfile=None):
|
|
191
|
+
"""certfile/keyfile serve HTTPS directly (sp_invoke_external_rest_endpoint only calls https:// URLs);
|
|
192
|
+
behind a TLS-terminating proxy or Cloud Run, leave them unset."""
|
|
193
|
+
srv = make_server(host, port, engine, token)
|
|
194
|
+
scheme = "http"
|
|
195
|
+
if certfile:
|
|
196
|
+
import ssl
|
|
197
|
+
ctx = ssl.SSLContext(ssl.PROTOCOL_TLS_SERVER)
|
|
198
|
+
ctx.load_cert_chain(certfile, keyfile)
|
|
199
|
+
accept = srv.get_request
|
|
200
|
+
|
|
201
|
+
def get_request():
|
|
202
|
+
# Handshake lazily, in the worker thread: a slow or failing client never blocks accept().
|
|
203
|
+
sock, addr = accept()
|
|
204
|
+
return ctx.wrap_socket(sock, server_side=True, do_handshake_on_connect=False), addr
|
|
205
|
+
srv.get_request = get_request
|
|
206
|
+
scheme = "https"
|
|
207
|
+
eng = srv.RequestHandlerClass.engine
|
|
208
|
+
print("sqljev gateway %s on %s://%s:%d (backend %s)" % (__version__, scheme, host, port, eng.cfg["backend"]),
|
|
209
|
+
file=sys.stderr)
|
|
210
|
+
try:
|
|
211
|
+
srv.serve_forever()
|
|
212
|
+
except KeyboardInterrupt:
|
|
213
|
+
pass
|
|
214
|
+
finally:
|
|
215
|
+
srv.server_close()
|
sqljev/spark.py
ADDED
|
@@ -0,0 +1,64 @@
|
|
|
1
|
+
"""Spark / Databricks: the jev functions as pandas (Arrow) UDFs usable from SQL.
|
|
2
|
+
|
|
3
|
+
import sqljev.spark
|
|
4
|
+
sqljev.spark.register(spark) # Laya in-process on each executor
|
|
5
|
+
# or: sqljev.spark.register(spark, backend="gateway", api_url="https://.../v1/eval",
|
|
6
|
+
# api_key=dbutils.secrets.get("sqljev", "token"))
|
|
7
|
+
|
|
8
|
+
SELECT * FROM tickets WHERE jev(to_json(struct(*)), 'the customer is angry')
|
|
9
|
+
|
|
10
|
+
Spark passes Arrow batches (spark.sql.execution.arrow.maxRecordsPerBatch, 10,000 rows by default) to each
|
|
11
|
+
UDF; each batch becomes one engine.call(). One engine per Python worker keeps the model loaded and the
|
|
12
|
+
answer cache warm across batches. On a GPU cluster set device="cuda".
|
|
13
|
+
"""
|
|
14
|
+
_ENGINES = {}
|
|
15
|
+
|
|
16
|
+
_RETURNS = {"jev": "boolean", "jev_prob": "double", "jev_score": "double", "jev_score_norm": "double",
|
|
17
|
+
"jev_choice": "string", "jev_confidence": "double", "jev_eval": "string"}
|
|
18
|
+
_ARITY = {"jev": 2, "jev_prob": 2, "jev_score": 3, "jev_score_norm": 3, "jev_choice": 3,
|
|
19
|
+
"jev_confidence": 4, "jev_eval": 4}
|
|
20
|
+
|
|
21
|
+
|
|
22
|
+
def _engine(settings):
|
|
23
|
+
from .core import Jev
|
|
24
|
+
key = tuple(sorted(settings.items()))
|
|
25
|
+
eng = _ENGINES.get(key)
|
|
26
|
+
if eng is None:
|
|
27
|
+
eng = _ENGINES[key] = Jev(**settings)
|
|
28
|
+
return eng
|
|
29
|
+
|
|
30
|
+
|
|
31
|
+
def _make(name, settings):
|
|
32
|
+
import json
|
|
33
|
+
import pandas as pd
|
|
34
|
+
from pyspark.sql.functions import pandas_udf
|
|
35
|
+
|
|
36
|
+
n = _ARITY[name]
|
|
37
|
+
|
|
38
|
+
def run(*cols):
|
|
39
|
+
calls = list(zip(*(c.tolist() for c in cols)))
|
|
40
|
+
out = _engine(settings).call(name, calls)
|
|
41
|
+
if name == "jev_eval":
|
|
42
|
+
out = [None if v is None else json.dumps(v) for v in out]
|
|
43
|
+
return pd.Series(out, dtype="object" if name in ("jev_choice", "jev_eval", "jev") else "float64")
|
|
44
|
+
|
|
45
|
+
# pandas_udf needs a fixed-arity signature with type hints.
|
|
46
|
+
if n == 2:
|
|
47
|
+
def f(a: pd.Series, b: pd.Series) -> pd.Series:
|
|
48
|
+
return run(a, b)
|
|
49
|
+
elif n == 3:
|
|
50
|
+
def f(a: pd.Series, b: pd.Series, c: pd.Series) -> pd.Series:
|
|
51
|
+
return run(a, b, c)
|
|
52
|
+
else:
|
|
53
|
+
def f(a: pd.Series, b: pd.Series, c: pd.Series, d: pd.Series) -> pd.Series:
|
|
54
|
+
return run(a, b, c, d)
|
|
55
|
+
return pandas_udf(f, _RETURNS[name])
|
|
56
|
+
|
|
57
|
+
|
|
58
|
+
def register(spark, prefix="", **settings):
|
|
59
|
+
"""Register jev, jev_prob, jev_score, jev_score_norm, jev_choice, jev_confidence, jev_eval for SQL.
|
|
60
|
+
Settings (backend, api_url, api_key, model, device, ...) are shipped to the executors."""
|
|
61
|
+
fns = {}
|
|
62
|
+
for name in _RETURNS:
|
|
63
|
+
fns[name] = spark.udf.register(prefix + name, _make(name, dict(settings)))
|
|
64
|
+
return fns
|