@relay-harness/coding-agent 1.0.2 → 1.0.3

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.
Files changed (75) hide show
  1. package/CHANGELOG.md +17 -6
  2. package/LICENSE +22 -0
  3. package/README.md +12 -13
  4. package/dist/bundle/chunks/{chunk-XSLAMYCK.js → chunk-5H4IRYOE.js} +555 -61
  5. package/dist/bundle/chunks/{chunk-H6S6DQTM.js → chunk-NS67SA75.js} +1 -1
  6. package/dist/bundle/chunks/{virtual-modules-FNVEVB65.js → virtual-modules-ZGB2J7DY.js} +1 -1
  7. package/dist/bundle/cli-runtime.js +1 -1
  8. package/dist/bundle/index.js +1 -1
  9. package/dist/bundle/rpc-entry.js +1 -1
  10. package/dist/core/settings-manager.d.ts +5 -3
  11. package/dist/core/settings-manager.d.ts.map +1 -1
  12. package/dist/core/settings-manager.js.map +1 -1
  13. package/dist/extensions/laya/assessment.d.ts +7 -1
  14. package/dist/extensions/laya/assessment.d.ts.map +1 -1
  15. package/dist/extensions/laya/assessment.js +21 -1
  16. package/dist/extensions/laya/assessment.js.map +1 -1
  17. package/dist/extensions/laya/docker.d.ts +19 -0
  18. package/dist/extensions/laya/docker.d.ts.map +1 -0
  19. package/dist/extensions/laya/docker.js +44 -0
  20. package/dist/extensions/laya/docker.js.map +1 -0
  21. package/dist/extensions/laya/index.d.ts +12 -3
  22. package/dist/extensions/laya/index.d.ts.map +1 -1
  23. package/dist/extensions/laya/index.js +331 -67
  24. package/dist/extensions/laya/index.js.map +1 -1
  25. package/dist/extensions/laya/learn.d.ts +67 -0
  26. package/dist/extensions/laya/learn.d.ts.map +1 -0
  27. package/dist/extensions/laya/learn.js +263 -0
  28. package/dist/extensions/laya/learn.js.map +1 -0
  29. package/dist/extensions/laya/memory.d.ts +48 -0
  30. package/dist/extensions/laya/memory.d.ts.map +1 -0
  31. package/dist/extensions/laya/memory.js +107 -0
  32. package/dist/extensions/laya/memory.js.map +1 -0
  33. package/dist/extensions/laya/model-manifest.d.ts +1 -1
  34. package/dist/extensions/laya/model-manifest.js +1 -1
  35. package/dist/extensions/laya/model-manifest.js.map +1 -1
  36. package/dist/extensions/laya/routers.d.ts +8 -0
  37. package/dist/extensions/laya/routers.d.ts.map +1 -1
  38. package/dist/extensions/laya/routers.js +2 -0
  39. package/dist/extensions/laya/routers.js.map +1 -1
  40. package/dist/extensions/laya/runtime.d.ts +39 -49
  41. package/dist/extensions/laya/runtime.d.ts.map +1 -1
  42. package/dist/extensions/laya/runtime.js +93 -145
  43. package/dist/extensions/laya/runtime.js.map +1 -1
  44. package/dist/extensions/laya/server.d.ts +39 -15
  45. package/dist/extensions/laya/server.d.ts.map +1 -1
  46. package/dist/extensions/laya/server.js +212 -78
  47. package/dist/extensions/laya/server.js.map +1 -1
  48. package/dist/extensions/laya/train-script.d.ts +8 -0
  49. package/dist/extensions/laya/train-script.d.ts.map +1 -0
  50. package/dist/extensions/laya/train-script.js +422 -0
  51. package/dist/extensions/laya/train-script.js.map +1 -0
  52. package/dist/extensions/laya/training.d.ts +176 -0
  53. package/dist/extensions/laya/training.d.ts.map +1 -0
  54. package/dist/extensions/laya/training.js +337 -0
  55. package/dist/extensions/laya/training.js.map +1 -0
  56. package/dist/modes/interactive/components/config-selector.d.ts.map +1 -1
  57. package/dist/modes/interactive/components/config-selector.js +3 -2
  58. package/dist/modes/interactive/components/config-selector.js.map +1 -1
  59. package/dist/modes/interactive/interactive-mode.d.ts.map +1 -1
  60. package/dist/modes/interactive/interactive-mode.js +4 -4
  61. package/dist/modes/interactive/interactive-mode.js.map +1 -1
  62. package/dist/utils/version-check.d.ts.map +1 -1
  63. package/dist/utils/version-check.js +10 -6
  64. package/dist/utils/version-check.js.map +1 -1
  65. package/docs/laya.md +59 -19
  66. package/docs/quickstart.md +19 -6
  67. package/docs/settings.md +4 -3
  68. package/examples/extensions/custom-provider-anthropic/package-lock.json +2 -2
  69. package/examples/extensions/custom-provider-anthropic/package.json +1 -1
  70. package/examples/extensions/custom-provider-gitlab-duo/package.json +1 -1
  71. package/examples/extensions/sandbox/package-lock.json +2 -2
  72. package/examples/extensions/sandbox/package.json +1 -1
  73. package/examples/extensions/with-deps/package-lock.json +2 -2
  74. package/examples/extensions/with-deps/package.json +1 -1
  75. package/package.json +15 -11
@@ -0,0 +1,422 @@
1
+ /**
2
+ * Training script for the routing model. It is built into the Laya Docker image as `train.py`, and
3
+ * `LayaTrainer` runs it in a training container. It is adapted from laya-trainer's
4
+ * `laya_ml.py`; the recipe and the hash split match, so its test split is the one the shipped model
5
+ * never saw.
6
+ */
7
+ export const TRAIN_SCRIPT = String.raw `"""Teaches Relay's Laya routing model the exercises of the training workspace.
8
+
9
+ Relay writes this script to the Laya home and runs it in the Laya Python environment. The recipe is
10
+ laya-trainer's: soft targets, a policy-gradient term on noisy logits and soft cross-entropy, then
11
+ temperature calibration on the validation split.
12
+
13
+ Training is incremental. It starts from the model that routes today and mixes the focus exercises
14
+ (the ones labeled from sessions) with a replay sample of the other exercises, so the model learns
15
+ the new tasks without forgetting the rest. Afterwards the new and the current model answer the
16
+ held-out test split, and Relay activates the new model only when it is not worse.
17
+
18
+ Progress goes to stderr as "PROGRESS <json>" lines. The result is one JSON object on stdout.
19
+ """
20
+ import os
21
+
22
+ os.environ.setdefault("USE_TF", "0")
23
+
24
+ import argparse
25
+ import gc
26
+ import hashlib
27
+ import json
28
+ import math
29
+ import random
30
+ import re
31
+ import sys
32
+ import time
33
+ from pathlib import Path
34
+
35
+ # Some libraries replace sys.stdout later; writing bytes keeps the output UTF-8 on every platform.
36
+ _OUT = getattr(sys.stdout, "buffer", None)
37
+ _ERR = getattr(sys.stderr, "buffer", None)
38
+
39
+
40
+ def _write(buf, fallback, text):
41
+ if buf is not None:
42
+ buf.write(text.encode("utf-8"))
43
+ buf.flush()
44
+ else:
45
+ fallback.write(text)
46
+ fallback.flush()
47
+
48
+
49
+ def emit(obj, code=0):
50
+ _write(_OUT, sys.stdout, json.dumps(obj, ensure_ascii=False) + "\n")
51
+ sys.exit(code)
52
+
53
+
54
+ def fail(error, message, **details):
55
+ emit({"ok": False, "error": error, "message": message, **details}, 1)
56
+
57
+
58
+ def progress(**fields):
59
+ _write(_ERR, sys.stderr, "PROGRESS " + json.dumps(fields) + "\n")
60
+
61
+
62
+ # ---------------------------------------------------------------- exercises
63
+ def norm(s):
64
+ return re.sub(r"\s+", " ", str(s).strip().lower())
65
+
66
+
67
+ def state_text(row):
68
+ s = row.get("state")
69
+ return s if isinstance(s, str) else json.dumps(s, ensure_ascii=False, sort_keys=True)
70
+
71
+
72
+ def split_of(row):
73
+ """The row's own split, or laya-trainer's fixed 80/10/10 split by text hash."""
74
+ if row.get("split") in ("train", "val", "test"):
75
+ return row["split"]
76
+ h = int(hashlib.sha1(norm(state_text(row)).encode()).hexdigest(), 16) % 100
77
+ return "test" if h < 10 else "val" if h < 20 else "train"
78
+
79
+
80
+ def read_jsonl(path):
81
+ rows = []
82
+ for line in Path(path).read_text(encoding="utf-8").splitlines():
83
+ if line.strip():
84
+ try:
85
+ rows.append(json.loads(line))
86
+ except json.JSONDecodeError:
87
+ continue
88
+ return rows
89
+
90
+
91
+ def question(spec):
92
+ q = {"type": spec["type"], "instructions": spec["instructions"]}
93
+ if spec["type"] != "noul" or spec.get("criteria"):
94
+ q["criteria"] = spec["criteria"]
95
+ return q
96
+
97
+
98
+ def target_for(spec, v, smooth=0.02):
99
+ t = spec["type"]
100
+ if t == "choice":
101
+ keys = list(spec["criteria"])
102
+ return [(1 - smooth) * (k == v) + smooth / len(keys) for k in keys]
103
+ if t == "noul":
104
+ y = 1.0 if v else 0.0
105
+ return [(1 - smooth) * (1 - y) + smooth / 2, (1 - smooth) * y + smooth / 2]
106
+ n = len(spec["criteria"]) # score: ordinal target, neighbors get some mass
107
+ p = [0.0] * n
108
+ p[v] = 0.8
109
+ near = [i for i in (v - 1, v + 1) if 0 <= i < n]
110
+ for i in near:
111
+ p[i] += 0.2 / len(near)
112
+ return p
113
+
114
+
115
+ def build_items(tok, cfg, rows, decs):
116
+ from laya.common import QTYPES, build_sequence, render_options
117
+
118
+ items, skipped = [], 0
119
+ for r in rows:
120
+ for qid, v in (r.get("expected") or {}).items():
121
+ spec = decs.get(qid)
122
+ if not spec:
123
+ continue
124
+ crit = spec.get("criteria", {} if spec["type"] == "choice" else None)
125
+ q = {"t": spec["type"], "ins": spec["instructions"], "crit": crit}
126
+ target = target_for(spec, v)
127
+ seq, markers = build_sequence(tok, r["state"], q, cfg["max_len"], cfg["head_max_len"])
128
+ if len(markers) != len(render_options({"t": spec["type"], "crit": crit})):
129
+ skipped += 1
130
+ continue
131
+ items.append({"ids": seq, "markers": markers, "qtype": QTYPES[spec["type"]], "target": target,
132
+ "label": target.index(max(target))})
133
+ return items, skipped
134
+
135
+
136
+ def collate(items, pad_id):
137
+ import torch
138
+
139
+ b, length = len(items), max(len(i["ids"]) for i in items)
140
+ k_max = max(len(i["markers"]) for i in items)
141
+ ids = torch.full((b, length), pad_id, dtype=torch.long)
142
+ att = torch.zeros((b, length), dtype=torch.long)
143
+ pos = torch.zeros((b, k_max), dtype=torch.long)
144
+ mask = torch.zeros((b, k_max), dtype=torch.bool)
145
+ tgt = torch.zeros((b, k_max))
146
+ for n, it in enumerate(items):
147
+ ids[n, :len(it["ids"])] = torch.tensor(it["ids"])
148
+ att[n, :len(it["ids"])] = 1
149
+ k = len(it["markers"])
150
+ pos[n, :k] = torch.tensor(it["markers"])
151
+ mask[n, :k] = True
152
+ tgt[n, :len(it["target"])] = torch.tensor(it["target"])
153
+ return ids, att, pos, mask, tgt, torch.tensor([i["qtype"] for i in items])
154
+
155
+
156
+ def fit_temperature(samples):
157
+ import torch
158
+ from laya.common import TEMP_MAX, TEMP_MIN
159
+
160
+ k_max = max(len(l) for l, _ in samples)
161
+ lg = torch.full((len(samples), k_max), -1e4)
162
+ tg = torch.zeros((len(samples), k_max))
163
+ for i, (l, t) in enumerate(samples):
164
+ lg[i, :len(l)] = torch.as_tensor(l)
165
+ tg[i, :len(t)] = torch.as_tensor(t, dtype=torch.float32)
166
+ lt = torch.zeros(1, requires_grad=True)
167
+ opt = torch.optim.LBFGS([lt], lr=0.1, max_iter=100)
168
+
169
+ def closure():
170
+ opt.zero_grad()
171
+ loss = -(tg * torch.log_softmax(lg / lt.exp(), -1)).sum(-1).mean()
172
+ loss.backward()
173
+ return loss
174
+
175
+ opt.step(closure)
176
+ return float(torch.clamp(lt.exp(), TEMP_MIN, TEMP_MAX).item())
177
+
178
+
179
+ def save_model(model, tok, cfg, path, meta):
180
+ from safetensors.torch import save_file
181
+
182
+ path.mkdir(parents=True, exist_ok=True)
183
+ save_file({k: v.detach().half().cpu().contiguous() for k, v in model.state_dict().items()},
184
+ str(path / "model.safetensors"))
185
+ model.encoder.config.save_pretrained(path / "encoder")
186
+ tok.save_pretrained(path / "tokenizer")
187
+ (path / "rl_agent_config.json").write_text(json.dumps(cfg, indent=2), encoding="utf-8")
188
+ (path / "laya_trainer_meta.json").write_text(json.dumps(meta, ensure_ascii=False, indent=2), encoding="utf-8")
189
+
190
+
191
+ def pick_device(requested):
192
+ import torch
193
+
194
+ if requested and requested != "auto":
195
+ return torch.device(requested)
196
+ if torch.cuda.is_available():
197
+ return torch.device("cuda")
198
+ if getattr(torch.backends, "mps", None) and torch.backends.mps.is_available():
199
+ return torch.device("mps")
200
+ return torch.device("cpu")
201
+
202
+
203
+ # ---------------------------------------------------------------- scoring
204
+ def predicted_label(spec, ans):
205
+ if spec["type"] == "choice":
206
+ return ans["choice"]
207
+ if spec["type"] == "noul":
208
+ return ans["noul"] >= 0.5
209
+ probs = ans["probabilities"]
210
+ return int(max(probs, key=lambda k: probs[k]))
211
+
212
+
213
+ def score(agent, rows, decs):
214
+ """Correct answers per question, over the rows that label it."""
215
+ per = {}
216
+ for qid, spec in decs.items():
217
+ sel = [r for r in rows if qid in (r.get("expected") or {})]
218
+ if not sel:
219
+ continue
220
+ outs = agent.predict_batch([r["state"] for r in sel], {qid: question(spec)}, batch_size=16)
221
+ correct = sum(1 for r, o in zip(sel, outs) if predicted_label(spec, o["answers"][qid]) == r["expected"][qid])
222
+ per[qid] = {"n": len(sel), "correct": correct}
223
+ return {"n": sum(m["n"] for m in per.values()), "correct": sum(m["correct"] for m in per.values()),
224
+ "per_question": per}
225
+
226
+
227
+ # ---------------------------------------------------------------- learn
228
+ def cmd_learn(a):
229
+ import torch
230
+ from safetensors.torch import load_file
231
+ from transformers import AutoTokenizer
232
+ from laya.common import build_model, proper_reward
233
+
234
+ t0 = time.time()
235
+ random.seed(a.seed)
236
+ torch.manual_seed(a.seed)
237
+ workspace = Path(a.workspace)
238
+ decs = json.loads((workspace / "decisions.json").read_text(encoding="utf-8"))["decisions"]
239
+ rows = [r for r in read_jsonl(workspace / "data" / "dataset.jsonl") if r.get("expected")]
240
+ focus = [r for r in rows if r.get("source") == a.focus_source and split_of(r) == "train"]
241
+ rest = [r for r in rows if r.get("source") != a.focus_source]
242
+ if not focus:
243
+ fail("no_focus", "There are no '%s' exercises to learn." % a.focus_source)
244
+ pool = [r for r in rest if split_of(r) == "train"]
245
+ val = [r for r in rest if split_of(r) == "val"][:a.val_rows]
246
+ test = [r for r in rows if split_of(r) == "test"]
247
+ replay_n = min(len(pool), max(a.replay_min, min(a.replay_max, a.replay_ratio * len(focus))))
248
+ replay = random.Random(a.seed).sample(pool, replay_n)
249
+
250
+ init = Path(a.init)
251
+ out = Path(a.out)
252
+ device = pick_device(a.device)
253
+ low_memory = a.low_memory == "on" or (
254
+ a.low_memory == "auto" and device.type == "cuda"
255
+ and torch.cuda.get_device_properties(device).total_memory < 12 * 1024 ** 3)
256
+ cfg = json.loads((init / "rl_agent_config.json").read_text(encoding="utf-8"))
257
+ cfg["gradient_checkpointing"] = True
258
+ tok = AutoTokenizer.from_pretrained(str(init / "tokenizer"))
259
+ train_items, skipped = build_items(tok, cfg, focus * a.repeat_focus + replay, decs)
260
+ val_items, _ = build_items(tok, cfg, val, decs)
261
+ if not train_items:
262
+ fail("no_items", "No exercise fits the model input.")
263
+
264
+ model = build_model(cfg, encoder_dir=str(init / "encoder"))
265
+ model.load_state_dict(load_file(str(init / "model.safetensors")), strict=True)
266
+ model.float()
267
+ model.encoder.gradient_checkpointing_enable(gradient_checkpointing_kwargs={"use_reentrant": False})
268
+ model.head_checkpointing = True
269
+ if low_memory: # freeze the embeddings and the lower half of the encoder: less optimizer memory
270
+ layers = max((int(m.group(1)) for n, _ in model.named_parameters()
271
+ if (m := re.search(r"layers\.(\d+)\.", n)) and "encoder." in n), default=-1) + 1
272
+ for n, p in model.named_parameters():
273
+ m = re.search(r"layers\.(\d+)\.", n)
274
+ if "encoder." in n and ("embeddings" in n or (m and int(m.group(1)) < layers // 2)):
275
+ p.requires_grad_(False)
276
+ model.to(device).train()
277
+ use_amp = device.type == "cuda" and torch.cuda.is_bf16_supported()
278
+
279
+ enc = [p for n, p in model.named_parameters() if "encoder." in n and p.requires_grad]
280
+ head = [p for n, p in model.named_parameters() if "encoder." not in n]
281
+ opt = torch.optim.AdamW([{"params": enc, "lr": a.lr_encoder}, {"params": head, "lr": a.lr_head}],
282
+ weight_decay=0.01)
283
+ updates = max(1, math.ceil(len(train_items) / a.micro_batch / a.grad_accum) * a.epochs)
284
+ sched = torch.optim.lr_scheduler.CosineAnnealingLR(opt, T_max=updates, eta_min=1e-6)
285
+
286
+ def forward(chunk):
287
+ ids, att, pos, mask, tgt, qt = [t.to(device) for t in collate(chunk, tok.pad_token_id)]
288
+ with torch.autocast(device.type, dtype=torch.bfloat16, enabled=use_amp):
289
+ logits, act = model(ids, att, pos, mask, qt)
290
+ return logits.float(), act, mask, tgt, qt
291
+
292
+ def evaluate(items):
293
+ model.eval()
294
+ loss = correct = 0
295
+ with torch.no_grad():
296
+ for s in range(0, len(items), a.micro_batch):
297
+ chunk = items[s:s + a.micro_batch]
298
+ lg, _, mask, tgt, _ = forward(chunk)
299
+ lp = torch.log_softmax(lg.masked_fill(~mask, -1e4), -1)
300
+ loss += -(tgt * lp).sum(-1).sum().item()
301
+ correct += (lp.argmax(-1).cpu() == torch.tensor([i["label"] for i in chunk])).sum().item()
302
+ model.train()
303
+ return loss / max(1, len(items)), correct / max(1, len(items))
304
+
305
+ steps = math.ceil(len(train_items) / a.micro_batch)
306
+ every = max(1, steps // 10)
307
+ history = []
308
+ progress(phase="train", epoch=0, epochs=a.epochs, done=0, device=device.type, low_memory=low_memory,
309
+ items=len(train_items))
310
+ try:
311
+ for ep in range(a.epochs):
312
+ random.Random(a.seed + ep).shuffle(train_items)
313
+ opt.zero_grad(set_to_none=True)
314
+ sigma = 0.4 + (0.1 - 0.4) * ep / max(1, a.epochs - 1)
315
+ total = batches = 0
316
+ started = time.time()
317
+ for s in range(0, len(train_items), a.micro_batch):
318
+ lg, act, mask, tgt, qt = forward(train_items[s:s + a.micro_batch])
319
+ k = mask.sum(-1, keepdim=True).float()
320
+ eps = torch.randn((4,) + lg.shape, device=device) * sigma * mask
321
+ eps = (eps - eps.sum(-1, keepdim=True) / k) * mask
322
+ noisy = lg.detach().unsqueeze(0) + eps
323
+ probs = torch.softmax(noisy.masked_fill(~mask, -1e4), -1)
324
+ with torch.no_grad():
325
+ reward = proper_reward(probs, tgt.unsqueeze(0), qt, mask, w_sph=0.75, w_rps=1.0)
326
+ adv = reward - reward.mean(0, keepdim=True)
327
+ adv = adv / (adv.std() + 1e-6)
328
+ logp = -(((noisy - lg.unsqueeze(0)) ** 2) * mask).sum(-1) / (2 * sigma ** 2)
329
+ loss_rl = -(adv * logp).mean()
330
+ loss_ce = -(tgt * torch.log_softmax(lg.masked_fill(~mask, -1e4), -1)).sum(-1).mean()
331
+ loss = (loss_rl + loss_ce + 0.0 * act.sum()) / a.grad_accum
332
+ loss.backward()
333
+ batches += 1
334
+ total += loss.item() * a.grad_accum
335
+ if batches % a.grad_accum == 0 or s + a.micro_batch >= len(train_items):
336
+ torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
337
+ opt.step()
338
+ sched.step()
339
+ opt.zero_grad(set_to_none=True)
340
+ if batches % every == 0 and batches < steps:
341
+ done = batches / steps
342
+ elapsed = time.time() - started
343
+ progress(phase="train", epoch=ep + 1, epochs=a.epochs, done=round(done, 3),
344
+ eta_s=round(elapsed / done * (1 - done) + (a.epochs - ep - 1) * elapsed / done))
345
+ val_loss, val_acc = evaluate(val_items) if val_items else (total / batches, None)
346
+ history.append({"epoch": ep + 1, "train_loss": round(total / batches, 4), "val_loss": round(val_loss, 4),
347
+ "val_acc": None if val_acc is None else round(val_acc, 4)})
348
+ progress(phase="train", epoch=ep + 1, epochs=a.epochs, done=1, val_acc=history[-1]["val_acc"])
349
+ except torch.cuda.OutOfMemoryError:
350
+ fail("out_of_memory", "The GPU ran out of memory while training.")
351
+
352
+ model.eval()
353
+ samples = [[], [], []]
354
+ with torch.no_grad():
355
+ for s in range(0, len(val_items), a.micro_batch):
356
+ chunk = val_items[s:s + a.micro_batch]
357
+ lg, _, _, _, _ = forward(chunk)
358
+ for i, it in enumerate(chunk):
359
+ samples[it["qtype"]].append((lg[i, :len(it["markers"])].cpu(), it["target"]))
360
+ previous = cfg.get("temperature") or [1.0, 1.0, 1.0]
361
+ cfg["temperature"] = [fit_temperature(g) if len(g) >= 10 else previous[i] for i, g in enumerate(samples)]
362
+ cfg.pop("temperature_by_options", None)
363
+ cfg["fine_tuned"] = True
364
+ meta = {"version": out.name, "init": str(init), "epochs": a.epochs, "history": history,
365
+ "temperatures": cfg["temperature"],
366
+ "rows": {"focus": len(focus), "replay": len(replay), "val": len(val), "test": len(test)},
367
+ "items": {"train": len(train_items), "val": len(val_items), "skipped": skipped},
368
+ "device": device.type, "low_memory": low_memory, "created": time.strftime("%Y-%m-%d %H:%M:%S")}
369
+ save_model(model, tok, cfg, out, meta)
370
+ del model, opt, sched
371
+ gc.collect()
372
+ if device.type == "cuda":
373
+ torch.cuda.empty_cache()
374
+
375
+ import laya
376
+
377
+ progress(phase="evaluate", model="new")
378
+ agent = laya.load(str(out), device=device.type)
379
+ candidate = score(agent, test, decs)
380
+ learned = score(agent, focus, decs)
381
+ del agent
382
+ gc.collect()
383
+ progress(phase="evaluate", model="current")
384
+ agent = laya.load(str(init), device=device.type)
385
+ current = score(agent, test, decs)
386
+ before = score(agent, focus, decs)
387
+ emit({"ok": True, "model": str(out), "seconds": round(time.time() - t0), "device": device.type,
388
+ "low_memory": low_memory, "rows": meta["rows"], "history": history,
389
+ "test": {"candidate": candidate, "current": current},
390
+ "focus": {"candidate": learned, "current": before}})
391
+
392
+
393
+ def main():
394
+ ap = argparse.ArgumentParser(description="Teach Relay's Laya routing model")
395
+ sub = ap.add_subparsers(dest="cmd", required=True)
396
+ p = sub.add_parser("learn")
397
+ p.add_argument("--workspace", required=True, help="folder with decisions.json and data/dataset.jsonl")
398
+ p.add_argument("--init", required=True, help="model to start from")
399
+ p.add_argument("--out", required=True, help="folder of the new model")
400
+ p.add_argument("--focus-source", default="session")
401
+ p.add_argument("--repeat-focus", type=int, default=4)
402
+ p.add_argument("--replay-min", type=int, default=100)
403
+ p.add_argument("--replay-ratio", type=int, default=4)
404
+ p.add_argument("--replay-max", type=int, default=600)
405
+ p.add_argument("--val-rows", type=int, default=100)
406
+ p.add_argument("--epochs", type=int, default=3)
407
+ p.add_argument("--micro-batch", type=int, default=2)
408
+ p.add_argument("--grad-accum", type=int, default=8)
409
+ p.add_argument("--lr-encoder", type=float, default=2e-5)
410
+ p.add_argument("--lr-head", type=float, default=1e-4)
411
+ p.add_argument("--low-memory", default="auto", choices=["auto", "on", "off"])
412
+ p.add_argument("--device", default="auto")
413
+ p.add_argument("--seed", type=int, default=7)
414
+ p.set_defaults(f=cmd_learn)
415
+ a = ap.parse_args()
416
+ a.f(a)
417
+
418
+
419
+ if __name__ == "__main__":
420
+ main()
421
+ `;
422
+ //# sourceMappingURL=train-script.js.map
@@ -0,0 +1 @@
1
+ {"version":3,"file":"train-script.js","sourceRoot":"","sources":["../../../src/extensions/laya/train-script.ts"],"names":[],"mappings":"AAAA;;;;;GAKG;AACH,MAAM,CAAC,MAAM,YAAY,GAAG,MAAM,CAAC,GAAG,CAAA;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;CA8ZrC,CAAC","sourcesContent":["/**\n * Training script for the routing model. It is built into the Laya Docker image as `train.py`, and\n * `LayaTrainer` runs it in a training container. It is adapted from laya-trainer's\n * `laya_ml.py`; the recipe and the hash split match, so its test split is the one the shipped model\n * never saw.\n */\nexport const TRAIN_SCRIPT = String.raw`\"\"\"Teaches Relay's Laya routing model the exercises of the training workspace.\n\nRelay writes this script to the Laya home and runs it in the Laya Python environment. The recipe is\nlaya-trainer's: soft targets, a policy-gradient term on noisy logits and soft cross-entropy, then\ntemperature calibration on the validation split.\n\nTraining is incremental. It starts from the model that routes today and mixes the focus exercises\n(the ones labeled from sessions) with a replay sample of the other exercises, so the model learns\nthe new tasks without forgetting the rest. Afterwards the new and the current model answer the\nheld-out test split, and Relay activates the new model only when it is not worse.\n\nProgress goes to stderr as \"PROGRESS <json>\" lines. The result is one JSON object on stdout.\n\"\"\"\nimport os\n\nos.environ.setdefault(\"USE_TF\", \"0\")\n\nimport argparse\nimport gc\nimport hashlib\nimport json\nimport math\nimport random\nimport re\nimport sys\nimport time\nfrom pathlib import Path\n\n# Some libraries replace sys.stdout later; writing bytes keeps the output UTF-8 on every platform.\n_OUT = getattr(sys.stdout, \"buffer\", None)\n_ERR = getattr(sys.stderr, \"buffer\", None)\n\n\ndef _write(buf, fallback, text):\n if buf is not None:\n buf.write(text.encode(\"utf-8\"))\n buf.flush()\n else:\n fallback.write(text)\n fallback.flush()\n\n\ndef emit(obj, code=0):\n _write(_OUT, sys.stdout, json.dumps(obj, ensure_ascii=False) + \"\\n\")\n sys.exit(code)\n\n\ndef fail(error, message, **details):\n emit({\"ok\": False, \"error\": error, \"message\": message, **details}, 1)\n\n\ndef progress(**fields):\n _write(_ERR, sys.stderr, \"PROGRESS \" + json.dumps(fields) + \"\\n\")\n\n\n# ---------------------------------------------------------------- exercises\ndef norm(s):\n return re.sub(r\"\\s+\", \" \", str(s).strip().lower())\n\n\ndef state_text(row):\n s = row.get(\"state\")\n return s if isinstance(s, str) else json.dumps(s, ensure_ascii=False, sort_keys=True)\n\n\ndef split_of(row):\n \"\"\"The row's own split, or laya-trainer's fixed 80/10/10 split by text hash.\"\"\"\n if row.get(\"split\") in (\"train\", \"val\", \"test\"):\n return row[\"split\"]\n h = int(hashlib.sha1(norm(state_text(row)).encode()).hexdigest(), 16) % 100\n return \"test\" if h < 10 else \"val\" if h < 20 else \"train\"\n\n\ndef read_jsonl(path):\n rows = []\n for line in Path(path).read_text(encoding=\"utf-8\").splitlines():\n if line.strip():\n try:\n rows.append(json.loads(line))\n except json.JSONDecodeError:\n continue\n return rows\n\n\ndef question(spec):\n q = {\"type\": spec[\"type\"], \"instructions\": spec[\"instructions\"]}\n if spec[\"type\"] != \"noul\" or spec.get(\"criteria\"):\n q[\"criteria\"] = spec[\"criteria\"]\n return q\n\n\ndef target_for(spec, v, smooth=0.02):\n t = spec[\"type\"]\n if t == \"choice\":\n keys = list(spec[\"criteria\"])\n return [(1 - smooth) * (k == v) + smooth / len(keys) for k in keys]\n if t == \"noul\":\n y = 1.0 if v else 0.0\n return [(1 - smooth) * (1 - y) + smooth / 2, (1 - smooth) * y + smooth / 2]\n n = len(spec[\"criteria\"]) # score: ordinal target, neighbors get some mass\n p = [0.0] * n\n p[v] = 0.8\n near = [i for i in (v - 1, v + 1) if 0 <= i < n]\n for i in near:\n p[i] += 0.2 / len(near)\n return p\n\n\ndef build_items(tok, cfg, rows, decs):\n from laya.common import QTYPES, build_sequence, render_options\n\n items, skipped = [], 0\n for r in rows:\n for qid, v in (r.get(\"expected\") or {}).items():\n spec = decs.get(qid)\n if not spec:\n continue\n crit = spec.get(\"criteria\", {} if spec[\"type\"] == \"choice\" else None)\n q = {\"t\": spec[\"type\"], \"ins\": spec[\"instructions\"], \"crit\": crit}\n target = target_for(spec, v)\n seq, markers = build_sequence(tok, r[\"state\"], q, cfg[\"max_len\"], cfg[\"head_max_len\"])\n if len(markers) != len(render_options({\"t\": spec[\"type\"], \"crit\": crit})):\n skipped += 1\n continue\n items.append({\"ids\": seq, \"markers\": markers, \"qtype\": QTYPES[spec[\"type\"]], \"target\": target,\n \"label\": target.index(max(target))})\n return items, skipped\n\n\ndef collate(items, pad_id):\n import torch\n\n b, length = len(items), max(len(i[\"ids\"]) for i in items)\n k_max = max(len(i[\"markers\"]) for i in items)\n ids = torch.full((b, length), pad_id, dtype=torch.long)\n att = torch.zeros((b, length), dtype=torch.long)\n pos = torch.zeros((b, k_max), dtype=torch.long)\n mask = torch.zeros((b, k_max), dtype=torch.bool)\n tgt = torch.zeros((b, k_max))\n for n, it in enumerate(items):\n ids[n, :len(it[\"ids\"])] = torch.tensor(it[\"ids\"])\n att[n, :len(it[\"ids\"])] = 1\n k = len(it[\"markers\"])\n pos[n, :k] = torch.tensor(it[\"markers\"])\n mask[n, :k] = True\n tgt[n, :len(it[\"target\"])] = torch.tensor(it[\"target\"])\n return ids, att, pos, mask, tgt, torch.tensor([i[\"qtype\"] for i in items])\n\n\ndef fit_temperature(samples):\n import torch\n from laya.common import TEMP_MAX, TEMP_MIN\n\n k_max = max(len(l) for l, _ in samples)\n lg = torch.full((len(samples), k_max), -1e4)\n tg = torch.zeros((len(samples), k_max))\n for i, (l, t) in enumerate(samples):\n lg[i, :len(l)] = torch.as_tensor(l)\n tg[i, :len(t)] = torch.as_tensor(t, dtype=torch.float32)\n lt = torch.zeros(1, requires_grad=True)\n opt = torch.optim.LBFGS([lt], lr=0.1, max_iter=100)\n\n def closure():\n opt.zero_grad()\n loss = -(tg * torch.log_softmax(lg / lt.exp(), -1)).sum(-1).mean()\n loss.backward()\n return loss\n\n opt.step(closure)\n return float(torch.clamp(lt.exp(), TEMP_MIN, TEMP_MAX).item())\n\n\ndef save_model(model, tok, cfg, path, meta):\n from safetensors.torch import save_file\n\n path.mkdir(parents=True, exist_ok=True)\n save_file({k: v.detach().half().cpu().contiguous() for k, v in model.state_dict().items()},\n str(path / \"model.safetensors\"))\n model.encoder.config.save_pretrained(path / \"encoder\")\n tok.save_pretrained(path / \"tokenizer\")\n (path / \"rl_agent_config.json\").write_text(json.dumps(cfg, indent=2), encoding=\"utf-8\")\n (path / \"laya_trainer_meta.json\").write_text(json.dumps(meta, ensure_ascii=False, indent=2), encoding=\"utf-8\")\n\n\ndef pick_device(requested):\n import torch\n\n if requested and requested != \"auto\":\n return torch.device(requested)\n if torch.cuda.is_available():\n return torch.device(\"cuda\")\n if getattr(torch.backends, \"mps\", None) and torch.backends.mps.is_available():\n return torch.device(\"mps\")\n return torch.device(\"cpu\")\n\n\n# ---------------------------------------------------------------- scoring\ndef predicted_label(spec, ans):\n if spec[\"type\"] == \"choice\":\n return ans[\"choice\"]\n if spec[\"type\"] == \"noul\":\n return ans[\"noul\"] >= 0.5\n probs = ans[\"probabilities\"]\n return int(max(probs, key=lambda k: probs[k]))\n\n\ndef score(agent, rows, decs):\n \"\"\"Correct answers per question, over the rows that label it.\"\"\"\n per = {}\n for qid, spec in decs.items():\n sel = [r for r in rows if qid in (r.get(\"expected\") or {})]\n if not sel:\n continue\n outs = agent.predict_batch([r[\"state\"] for r in sel], {qid: question(spec)}, batch_size=16)\n correct = sum(1 for r, o in zip(sel, outs) if predicted_label(spec, o[\"answers\"][qid]) == r[\"expected\"][qid])\n per[qid] = {\"n\": len(sel), \"correct\": correct}\n return {\"n\": sum(m[\"n\"] for m in per.values()), \"correct\": sum(m[\"correct\"] for m in per.values()),\n \"per_question\": per}\n\n\n# ---------------------------------------------------------------- learn\ndef cmd_learn(a):\n import torch\n from safetensors.torch import load_file\n from transformers import AutoTokenizer\n from laya.common import build_model, proper_reward\n\n t0 = time.time()\n random.seed(a.seed)\n torch.manual_seed(a.seed)\n workspace = Path(a.workspace)\n decs = json.loads((workspace / \"decisions.json\").read_text(encoding=\"utf-8\"))[\"decisions\"]\n rows = [r for r in read_jsonl(workspace / \"data\" / \"dataset.jsonl\") if r.get(\"expected\")]\n focus = [r for r in rows if r.get(\"source\") == a.focus_source and split_of(r) == \"train\"]\n rest = [r for r in rows if r.get(\"source\") != a.focus_source]\n if not focus:\n fail(\"no_focus\", \"There are no '%s' exercises to learn.\" % a.focus_source)\n pool = [r for r in rest if split_of(r) == \"train\"]\n val = [r for r in rest if split_of(r) == \"val\"][:a.val_rows]\n test = [r for r in rows if split_of(r) == \"test\"]\n replay_n = min(len(pool), max(a.replay_min, min(a.replay_max, a.replay_ratio * len(focus))))\n replay = random.Random(a.seed).sample(pool, replay_n)\n\n init = Path(a.init)\n out = Path(a.out)\n device = pick_device(a.device)\n low_memory = a.low_memory == \"on\" or (\n a.low_memory == \"auto\" and device.type == \"cuda\"\n and torch.cuda.get_device_properties(device).total_memory < 12 * 1024 ** 3)\n cfg = json.loads((init / \"rl_agent_config.json\").read_text(encoding=\"utf-8\"))\n cfg[\"gradient_checkpointing\"] = True\n tok = AutoTokenizer.from_pretrained(str(init / \"tokenizer\"))\n train_items, skipped = build_items(tok, cfg, focus * a.repeat_focus + replay, decs)\n val_items, _ = build_items(tok, cfg, val, decs)\n if not train_items:\n fail(\"no_items\", \"No exercise fits the model input.\")\n\n model = build_model(cfg, encoder_dir=str(init / \"encoder\"))\n model.load_state_dict(load_file(str(init / \"model.safetensors\")), strict=True)\n model.float()\n model.encoder.gradient_checkpointing_enable(gradient_checkpointing_kwargs={\"use_reentrant\": False})\n model.head_checkpointing = True\n if low_memory: # freeze the embeddings and the lower half of the encoder: less optimizer memory\n layers = max((int(m.group(1)) for n, _ in model.named_parameters()\n if (m := re.search(r\"layers\\.(\\d+)\\.\", n)) and \"encoder.\" in n), default=-1) + 1\n for n, p in model.named_parameters():\n m = re.search(r\"layers\\.(\\d+)\\.\", n)\n if \"encoder.\" in n and (\"embeddings\" in n or (m and int(m.group(1)) < layers // 2)):\n p.requires_grad_(False)\n model.to(device).train()\n use_amp = device.type == \"cuda\" and torch.cuda.is_bf16_supported()\n\n enc = [p for n, p in model.named_parameters() if \"encoder.\" in n and p.requires_grad]\n head = [p for n, p in model.named_parameters() if \"encoder.\" not in n]\n opt = torch.optim.AdamW([{\"params\": enc, \"lr\": a.lr_encoder}, {\"params\": head, \"lr\": a.lr_head}],\n weight_decay=0.01)\n updates = max(1, math.ceil(len(train_items) / a.micro_batch / a.grad_accum) * a.epochs)\n sched = torch.optim.lr_scheduler.CosineAnnealingLR(opt, T_max=updates, eta_min=1e-6)\n\n def forward(chunk):\n ids, att, pos, mask, tgt, qt = [t.to(device) for t in collate(chunk, tok.pad_token_id)]\n with torch.autocast(device.type, dtype=torch.bfloat16, enabled=use_amp):\n logits, act = model(ids, att, pos, mask, qt)\n return logits.float(), act, mask, tgt, qt\n\n def evaluate(items):\n model.eval()\n loss = correct = 0\n with torch.no_grad():\n for s in range(0, len(items), a.micro_batch):\n chunk = items[s:s + a.micro_batch]\n lg, _, mask, tgt, _ = forward(chunk)\n lp = torch.log_softmax(lg.masked_fill(~mask, -1e4), -1)\n loss += -(tgt * lp).sum(-1).sum().item()\n correct += (lp.argmax(-1).cpu() == torch.tensor([i[\"label\"] for i in chunk])).sum().item()\n model.train()\n return loss / max(1, len(items)), correct / max(1, len(items))\n\n steps = math.ceil(len(train_items) / a.micro_batch)\n every = max(1, steps // 10)\n history = []\n progress(phase=\"train\", epoch=0, epochs=a.epochs, done=0, device=device.type, low_memory=low_memory,\n items=len(train_items))\n try:\n for ep in range(a.epochs):\n random.Random(a.seed + ep).shuffle(train_items)\n opt.zero_grad(set_to_none=True)\n sigma = 0.4 + (0.1 - 0.4) * ep / max(1, a.epochs - 1)\n total = batches = 0\n started = time.time()\n for s in range(0, len(train_items), a.micro_batch):\n lg, act, mask, tgt, qt = forward(train_items[s:s + a.micro_batch])\n k = mask.sum(-1, keepdim=True).float()\n eps = torch.randn((4,) + lg.shape, device=device) * sigma * mask\n eps = (eps - eps.sum(-1, keepdim=True) / k) * mask\n noisy = lg.detach().unsqueeze(0) + eps\n probs = torch.softmax(noisy.masked_fill(~mask, -1e4), -1)\n with torch.no_grad():\n reward = proper_reward(probs, tgt.unsqueeze(0), qt, mask, w_sph=0.75, w_rps=1.0)\n adv = reward - reward.mean(0, keepdim=True)\n adv = adv / (adv.std() + 1e-6)\n logp = -(((noisy - lg.unsqueeze(0)) ** 2) * mask).sum(-1) / (2 * sigma ** 2)\n loss_rl = -(adv * logp).mean()\n loss_ce = -(tgt * torch.log_softmax(lg.masked_fill(~mask, -1e4), -1)).sum(-1).mean()\n loss = (loss_rl + loss_ce + 0.0 * act.sum()) / a.grad_accum\n loss.backward()\n batches += 1\n total += loss.item() * a.grad_accum\n if batches % a.grad_accum == 0 or s + a.micro_batch >= len(train_items):\n torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)\n opt.step()\n sched.step()\n opt.zero_grad(set_to_none=True)\n if batches % every == 0 and batches < steps:\n done = batches / steps\n elapsed = time.time() - started\n progress(phase=\"train\", epoch=ep + 1, epochs=a.epochs, done=round(done, 3),\n eta_s=round(elapsed / done * (1 - done) + (a.epochs - ep - 1) * elapsed / done))\n val_loss, val_acc = evaluate(val_items) if val_items else (total / batches, None)\n history.append({\"epoch\": ep + 1, \"train_loss\": round(total / batches, 4), \"val_loss\": round(val_loss, 4),\n \"val_acc\": None if val_acc is None else round(val_acc, 4)})\n progress(phase=\"train\", epoch=ep + 1, epochs=a.epochs, done=1, val_acc=history[-1][\"val_acc\"])\n except torch.cuda.OutOfMemoryError:\n fail(\"out_of_memory\", \"The GPU ran out of memory while training.\")\n\n model.eval()\n samples = [[], [], []]\n with torch.no_grad():\n for s in range(0, len(val_items), a.micro_batch):\n chunk = val_items[s:s + a.micro_batch]\n lg, _, _, _, _ = forward(chunk)\n for i, it in enumerate(chunk):\n samples[it[\"qtype\"]].append((lg[i, :len(it[\"markers\"])].cpu(), it[\"target\"]))\n previous = cfg.get(\"temperature\") or [1.0, 1.0, 1.0]\n cfg[\"temperature\"] = [fit_temperature(g) if len(g) >= 10 else previous[i] for i, g in enumerate(samples)]\n cfg.pop(\"temperature_by_options\", None)\n cfg[\"fine_tuned\"] = True\n meta = {\"version\": out.name, \"init\": str(init), \"epochs\": a.epochs, \"history\": history,\n \"temperatures\": cfg[\"temperature\"],\n \"rows\": {\"focus\": len(focus), \"replay\": len(replay), \"val\": len(val), \"test\": len(test)},\n \"items\": {\"train\": len(train_items), \"val\": len(val_items), \"skipped\": skipped},\n \"device\": device.type, \"low_memory\": low_memory, \"created\": time.strftime(\"%Y-%m-%d %H:%M:%S\")}\n save_model(model, tok, cfg, out, meta)\n del model, opt, sched\n gc.collect()\n if device.type == \"cuda\":\n torch.cuda.empty_cache()\n\n import laya\n\n progress(phase=\"evaluate\", model=\"new\")\n agent = laya.load(str(out), device=device.type)\n candidate = score(agent, test, decs)\n learned = score(agent, focus, decs)\n del agent\n gc.collect()\n progress(phase=\"evaluate\", model=\"current\")\n agent = laya.load(str(init), device=device.type)\n current = score(agent, test, decs)\n before = score(agent, focus, decs)\n emit({\"ok\": True, \"model\": str(out), \"seconds\": round(time.time() - t0), \"device\": device.type,\n \"low_memory\": low_memory, \"rows\": meta[\"rows\"], \"history\": history,\n \"test\": {\"candidate\": candidate, \"current\": current},\n \"focus\": {\"candidate\": learned, \"current\": before}})\n\n\ndef main():\n ap = argparse.ArgumentParser(description=\"Teach Relay's Laya routing model\")\n sub = ap.add_subparsers(dest=\"cmd\", required=True)\n p = sub.add_parser(\"learn\")\n p.add_argument(\"--workspace\", required=True, help=\"folder with decisions.json and data/dataset.jsonl\")\n p.add_argument(\"--init\", required=True, help=\"model to start from\")\n p.add_argument(\"--out\", required=True, help=\"folder of the new model\")\n p.add_argument(\"--focus-source\", default=\"session\")\n p.add_argument(\"--repeat-focus\", type=int, default=4)\n p.add_argument(\"--replay-min\", type=int, default=100)\n p.add_argument(\"--replay-ratio\", type=int, default=4)\n p.add_argument(\"--replay-max\", type=int, default=600)\n p.add_argument(\"--val-rows\", type=int, default=100)\n p.add_argument(\"--epochs\", type=int, default=3)\n p.add_argument(\"--micro-batch\", type=int, default=2)\n p.add_argument(\"--grad-accum\", type=int, default=8)\n p.add_argument(\"--lr-encoder\", type=float, default=2e-5)\n p.add_argument(\"--lr-head\", type=float, default=1e-4)\n p.add_argument(\"--low-memory\", default=\"auto\", choices=[\"auto\", \"on\", \"off\"])\n p.add_argument(\"--device\", default=\"auto\")\n p.add_argument(\"--seed\", type=int, default=7)\n p.set_defaults(f=cmd_learn)\n a = ap.parse_args()\n a.f(a)\n\n\nif __name__ == \"__main__\":\n main()\n`;\n"]}
@@ -0,0 +1,176 @@
1
+ import { type DockerRun } from "./docker.ts";
2
+ import { type LayaModelManifest } from "./runtime.ts";
3
+ /**
4
+ * Training of the routing model, in a Docker container. Exercises live in a workspace in the Laya
5
+ * home on the host: the synthetic seed the shipped model was trained on, plus the tasks labeled from
6
+ * sessions (`/laya learn`). The container reads the workspace and writes the new model to the models
7
+ * volume. Training starts from the model that routes today, and the new model replaces it only when
8
+ * it answers the held-out test split about as well.
9
+ */
10
+ /** One exercise in laya-trainer's dataset format. */
11
+ export interface TrainingRow {
12
+ id: string;
13
+ state: {
14
+ request: string;
15
+ };
16
+ expected: Record<string, string | number | boolean>;
17
+ source: string;
18
+ split?: "train" | "val" | "test";
19
+ /** Session the task came from. */
20
+ session?: string;
21
+ label_source?: string;
22
+ note?: string;
23
+ /** What to know when doing a similar task. Shown to the model, never trained on. */
24
+ lesson?: string;
25
+ created?: string;
26
+ }
27
+ export interface TrainingWorkspace {
28
+ dir: string;
29
+ decisions: string;
30
+ dataset: string;
31
+ registry: string;
32
+ }
33
+ /** Exercises labeled from sessions. Training focuses on them; the rest is replayed. */
34
+ export declare const SESSION_SOURCE = "session";
35
+ /** The seed size of the shipped model: the same rows, so the same held-out test split. */
36
+ export declare const SEED_ROWS = 1100;
37
+ /** A new model is activated unless it answers more than this fraction fewer test questions. */
38
+ export declare const REGRESSION_TOLERANCE = 0.01;
39
+ export declare function trainingWorkspace(home: string): TrainingWorkspace;
40
+ /** Writes the current questions and, on first use, the seed exercises. */
41
+ export declare function ensureWorkspace(workspace: TrainingWorkspace): void;
42
+ export declare function readRows(path: string): TrainingRow[];
43
+ export declare function writeRows(path: string, rows: readonly TrainingRow[]): void;
44
+ /** Requests that differ only in case and spacing are the same exercise. */
45
+ export declare function requestKey(request: string): string;
46
+ /**
47
+ * Adds rows, replacing the labels of an exercise with the same request: a later session labels it
48
+ * with newer evidence.
49
+ */
50
+ export declare function mergeRows(existing: readonly TrainingRow[], incoming: ReadonlyArray<Omit<TrainingRow, "id">>): {
51
+ rows: TrainingRow[];
52
+ added: number;
53
+ updated: number;
54
+ };
55
+ /** Correct answers on a set of exercises, as the training script reports them. */
56
+ export interface Score {
57
+ n: number;
58
+ correct: number;
59
+ per_question: Record<string, {
60
+ n: number;
61
+ correct: number;
62
+ }>;
63
+ }
64
+ export interface TrainedModel {
65
+ name: string;
66
+ createdAt: string;
67
+ /** Model the training started from. */
68
+ basedOn: string;
69
+ sessionTasks: number;
70
+ /** Test split: the new model and the one it started from. */
71
+ test: {
72
+ candidate: Score;
73
+ current: Score;
74
+ };
75
+ /** Session exercises: answers the new model and the one it started from get right. */
76
+ session: {
77
+ candidate: Score;
78
+ current: Score;
79
+ };
80
+ seconds: number;
81
+ device: string;
82
+ }
83
+ export interface ModelRegistry {
84
+ /** Trained model that routes requests. Absent: the shipped model. */
85
+ active?: string;
86
+ models: TrainedModel[];
87
+ }
88
+ export declare function readRegistry(workspace: TrainingWorkspace): ModelRegistry;
89
+ export declare function writeRegistry(workspace: TrainingWorkspace, registry: ModelRegistry): void;
90
+ export declare const accuracy: (score: Score) => number;
91
+ /** Whether the new model answers the test split well enough to replace the current one. */
92
+ export declare function passesGate(test: TrainedModel["test"]): boolean;
93
+ /** Keeps the newest inactive models and the active one; returns the names to delete. */
94
+ export declare function pruneModels(registry: ModelRegistry): {
95
+ registry: ModelRegistry;
96
+ removed: string[];
97
+ };
98
+ export interface TrainingProgress {
99
+ phase: "train" | "evaluate";
100
+ epoch?: number;
101
+ epochs?: number;
102
+ done?: number;
103
+ eta_s?: number;
104
+ device?: string;
105
+ model?: "new" | "current";
106
+ }
107
+ export declare function describeProgress(progress: TrainingProgress): string;
108
+ export interface TrainingOutcome {
109
+ model: TrainedModel;
110
+ /** Whether the new model now routes requests. */
111
+ activated: boolean;
112
+ /** Model that routed before. */
113
+ previous: string;
114
+ }
115
+ export interface TrainingListener {
116
+ progress(progress: TrainingProgress): void;
117
+ finished(result: {
118
+ outcome: TrainingOutcome;
119
+ } | {
120
+ error: Error;
121
+ }): void;
122
+ }
123
+ export interface LayaTrainerOptions {
124
+ home: string;
125
+ manifest: LayaModelManifest;
126
+ docker: DockerRun;
127
+ /** Image the training container runs; the server's image. */
128
+ image: () => Promise<string>;
129
+ /** Whether the training container can use an NVIDIA GPU. */
130
+ gpu: () => Promise<boolean>;
131
+ }
132
+ /** Whether a container of the image sees an NVIDIA GPU: needs the CUDA image and Docker GPU support. */
133
+ export declare function dockerHasGpu(docker: DockerRun, image: string): Promise<boolean>;
134
+ /**
135
+ * Runs one training at a time. It belongs to the process, not to a session runtime, so training
136
+ * goes on across `/new` and `/resume`; `listener` is the latest runtime that wants its events.
137
+ */
138
+ export declare class LayaTrainer {
139
+ readonly workspace: TrainingWorkspace;
140
+ private readonly options;
141
+ private abort;
142
+ private job;
143
+ /** Latest progress of the running training. */
144
+ progress: TrainingProgress | undefined;
145
+ listener: TrainingListener | undefined;
146
+ constructor(options: LayaTrainerOptions);
147
+ get running(): boolean;
148
+ /** The model that routes requests, with its path inside the containers: the active trained model, or the shipped one. */
149
+ activeModel(): {
150
+ name: string;
151
+ dir: string;
152
+ };
153
+ /** Selects the model that routes requests; the shipped model's version selects it. */
154
+ use(name: string): void;
155
+ /** Exercises in the workspace, by source. */
156
+ counts(): Record<string, number>;
157
+ /** Adds labeled session tasks to the workspace. */
158
+ addExercises(rows: ReadonlyArray<Omit<TrainingRow, "id">>): {
159
+ added: number;
160
+ updated: number;
161
+ };
162
+ /**
163
+ * Trains a new model from the active one and activates it when it passes the test gate. The
164
+ * listener hears the result too, so a caller that does not wait only needs to catch.
165
+ */
166
+ train(): Promise<TrainingOutcome>;
167
+ /** Stops a running training and removes its container. */
168
+ stop(): void;
169
+ private run;
170
+ /** Deletes trained models from the volume. */
171
+ private removeModels;
172
+ private runScript;
173
+ }
174
+ /** The trainer of a Laya home, shared by every session runtime of the process. */
175
+ export declare function sharedTrainer(options: LayaTrainerOptions): LayaTrainer;
176
+ //# sourceMappingURL=training.d.ts.map
@@ -0,0 +1 @@
1
+ {"version":3,"file":"training.d.ts","sourceRoot":"","sources":["../../../src/extensions/laya/training.ts"],"names":[],"mappings":"AAEA,OAAO,EAAE,KAAK,SAAS,EAAe,MAAM,aAAa,CAAC;AAE1D,OAAO,EAIN,KAAK,iBAAiB,EAEtB,MAAM,cAAc,CAAC;AAGtB;;;;;;GAMG;AAEH,qDAAqD;AACrD,MAAM,WAAW,WAAW;IAC3B,EAAE,EAAE,MAAM,CAAC;IACX,KAAK,EAAE;QAAE,OAAO,EAAE,MAAM,CAAA;KAAE,CAAC;IAC3B,QAAQ,EAAE,MAAM,CAAC,MAAM,EAAE,MAAM,GAAG,MAAM,GAAG,OAAO,CAAC,CAAC;IACpD,MAAM,EAAE,MAAM,CAAC;IACf,KAAK,CAAC,EAAE,OAAO,GAAG,KAAK,GAAG,MAAM,CAAC;IACjC,kCAAkC;IAClC,OAAO,CAAC,EAAE,MAAM,CAAC;IACjB,YAAY,CAAC,EAAE,MAAM,CAAC;IACtB,IAAI,CAAC,EAAE,MAAM,CAAC;IACd,oFAAoF;IACpF,MAAM,CAAC,EAAE,MAAM,CAAC;IAChB,OAAO,CAAC,EAAE,MAAM,CAAC;CACjB;AAED,MAAM,WAAW,iBAAiB;IACjC,GAAG,EAAE,MAAM,CAAC;IACZ,SAAS,EAAE,MAAM,CAAC;IAClB,OAAO,EAAE,MAAM,CAAC;IAChB,QAAQ,EAAE,MAAM,CAAC;CACjB;AAED,uFAAuF;AACvF,eAAO,MAAM,cAAc,YAAY,CAAC;AACxC,0FAA0F;AAC1F,eAAO,MAAM,SAAS,OAAO,CAAC;AAC9B,+FAA+F;AAC/F,eAAO,MAAM,oBAAoB,OAAO,CAAC;AAIzC,wBAAgB,iBAAiB,CAAC,IAAI,EAAE,MAAM,GAAG,iBAAiB,CAQjE;AAED,0EAA0E;AAC1E,wBAAgB,eAAe,CAAC,SAAS,EAAE,iBAAiB,GAAG,IAAI,CAOlE;AAED,wBAAgB,QAAQ,CAAC,IAAI,EAAE,MAAM,GAAG,WAAW,EAAE,CAYpD;AAED,wBAAgB,SAAS,CAAC,IAAI,EAAE,MAAM,EAAE,IAAI,EAAE,SAAS,WAAW,EAAE,GAAG,IAAI,CAG1E;AAED,2EAA2E;AAC3E,wBAAgB,UAAU,CAAC,OAAO,EAAE,MAAM,GAAG,MAAM,CAElD;AAED;;;GAGG;AACH,wBAAgB,SAAS,CACxB,QAAQ,EAAE,SAAS,WAAW,EAAE,EAChC,QAAQ,EAAE,aAAa,CAAC,IAAI,CAAC,WAAW,EAAE,IAAI,CAAC,CAAC,GAC9C;IAAE,IAAI,EAAE,WAAW,EAAE,CAAC;IAAC,KAAK,EAAE,MAAM,CAAC;IAAC,OAAO,EAAE,MAAM,CAAA;CAAE,CAmBzD;AAED,kFAAkF;AAClF,MAAM,WAAW,KAAK;IACrB,CAAC,EAAE,MAAM,CAAC;IACV,OAAO,EAAE,MAAM,CAAC;IAChB,YAAY,EAAE,MAAM,CAAC,MAAM,EAAE;QAAE,CAAC,EAAE,MAAM,CAAC;QAAC,OAAO,EAAE,MAAM,CAAA;KAAE,CAAC,CAAC;CAC7D;AAED,MAAM,WAAW,YAAY;IAC5B,IAAI,EAAE,MAAM,CAAC;IACb,SAAS,EAAE,MAAM,CAAC;IAClB,uCAAuC;IACvC,OAAO,EAAE,MAAM,CAAC;IAChB,YAAY,EAAE,MAAM,CAAC;IACrB,6DAA6D;IAC7D,IAAI,EAAE;QAAE,SAAS,EAAE,KAAK,CAAC;QAAC,OAAO,EAAE,KAAK,CAAA;KAAE,CAAC;IAC3C,sFAAsF;IACtF,OAAO,EAAE;QAAE,SAAS,EAAE,KAAK,CAAC;QAAC,OAAO,EAAE,KAAK,CAAA;KAAE,CAAC;IAC9C,OAAO,EAAE,MAAM,CAAC;IAChB,MAAM,EAAE,MAAM,CAAC;CACf;AAED,MAAM,WAAW,aAAa;IAC7B,qEAAqE;IACrE,MAAM,CAAC,EAAE,MAAM,CAAC;IAChB,MAAM,EAAE,YAAY,EAAE,CAAC;CACvB;AAED,wBAAgB,YAAY,CAAC,SAAS,EAAE,iBAAiB,GAAG,aAAa,CAOxE;AAED,wBAAgB,aAAa,CAAC,SAAS,EAAE,iBAAiB,EAAE,QAAQ,EAAE,aAAa,GAAG,IAAI,CAGzF;AAED,eAAO,MAAM,QAAQ,UAAW,KAAK,KAAG,MAAqD,CAAC;AAE9F,2FAA2F;AAC3F,wBAAgB,UAAU,CAAC,IAAI,EAAE,YAAY,CAAC,MAAM,CAAC,GAAG,OAAO,CAE9D;AAED,wFAAwF;AACxF,wBAAgB,WAAW,CAAC,QAAQ,EAAE,aAAa,GAAG;IAAE,QAAQ,EAAE,aAAa,CAAC;IAAC,OAAO,EAAE,MAAM,EAAE,CAAA;CAAE,CAQnG;AAED,MAAM,WAAW,gBAAgB;IAChC,KAAK,EAAE,OAAO,GAAG,UAAU,CAAC;IAC5B,KAAK,CAAC,EAAE,MAAM,CAAC;IACf,MAAM,CAAC,EAAE,MAAM,CAAC;IAChB,IAAI,CAAC,EAAE,MAAM,CAAC;IACd,KAAK,CAAC,EAAE,MAAM,CAAC;IACf,MAAM,CAAC,EAAE,MAAM,CAAC;IAChB,KAAK,CAAC,EAAE,KAAK,GAAG,SAAS,CAAC;CAC1B;AAED,wBAAgB,gBAAgB,CAAC,QAAQ,EAAE,gBAAgB,GAAG,MAAM,CASnE;AAaD,MAAM,WAAW,eAAe;IAC/B,KAAK,EAAE,YAAY,CAAC;IACpB,iDAAiD;IACjD,SAAS,EAAE,OAAO,CAAC;IACnB,gCAAgC;IAChC,QAAQ,EAAE,MAAM,CAAC;CACjB;AAED,MAAM,WAAW,gBAAgB;IAChC,QAAQ,CAAC,QAAQ,EAAE,gBAAgB,GAAG,IAAI,CAAC;IAC3C,QAAQ,CAAC,MAAM,EAAE;QAAE,OAAO,EAAE,eAAe,CAAA;KAAE,GAAG;QAAE,KAAK,EAAE,KAAK,CAAA;KAAE,GAAG,IAAI,CAAC;CACxE;AAED,MAAM,WAAW,kBAAkB;IAClC,IAAI,EAAE,MAAM,CAAC;IACb,QAAQ,EAAE,iBAAiB,CAAC;IAC5B,MAAM,EAAE,SAAS,CAAC;IAClB,6DAA6D;IAC7D,KAAK,EAAE,MAAM,OAAO,CAAC,MAAM,CAAC,CAAC;IAC7B,4DAA4D;IAC5D,GAAG,EAAE,MAAM,OAAO,CAAC,OAAO,CAAC,CAAC;CAC5B;AAkBD,wGAAwG;AACxG,wBAAsB,YAAY,CAAC,MAAM,EAAE,SAAS,EAAE,KAAK,EAAE,MAAM,GAAG,OAAO,CAAC,OAAO,CAAC,CAMrF;AAED;;;GAGG;AACH,qBAAa,WAAW;IACvB,QAAQ,CAAC,SAAS,EAAE,iBAAiB,CAAC;IACtC,OAAO,CAAC,QAAQ,CAAC,OAAO,CAAqB;IAC7C,OAAO,CAAC,KAAK,CAA8B;IAC3C,OAAO,CAAC,GAAG,CAAuC;IAClD,+CAA+C;IAC/C,QAAQ,EAAE,gBAAgB,GAAG,SAAS,CAAC;IACvC,QAAQ,EAAE,gBAAgB,GAAG,SAAS,CAAC;IAEvC,YAAY,OAAO,EAAE,kBAAkB,EAGtC;IAED,IAAI,OAAO,IAAI,OAAO,CAErB;IAED,yHAAyH;IACzH,WAAW,IAAI;QAAE,IAAI,EAAE,MAAM,CAAC;QAAC,GAAG,EAAE,MAAM,CAAA;KAAE,CAO3C;IAED,sFAAsF;IACtF,GAAG,CAAC,IAAI,EAAE,MAAM,GAAG,IAAI,CAYtB;IAED,6CAA6C;IAC7C,MAAM,IAAI,MAAM,CAAC,MAAM,EAAE,MAAM,CAAC,CAI/B;IAED,mDAAmD;IACnD,YAAY,CAAC,IAAI,EAAE,aAAa,CAAC,IAAI,CAAC,WAAW,EAAE,IAAI,CAAC,CAAC,GAAG;QAAE,KAAK,EAAE,MAAM,CAAC;QAAC,OAAO,EAAE,MAAM,CAAA;KAAE,CAK7F;IAED;;;OAGG;IACH,KAAK,IAAI,OAAO,CAAC,eAAe,CAAC,CAqBhC;IAED,0DAA0D;IAC1D,IAAI,IAAI,IAAI,CAIX;YAEa,GAAG;IAoEjB,8CAA8C;YAChC,YAAY;YAYZ,SAAS;CA4BvB;AAID,kFAAkF;AAClF,wBAAgB,aAAa,CAAC,OAAO,EAAE,kBAAkB,GAAG,WAAW,CAOtE"}