Quazim0t0 commited on
Commit
99509b7
·
verified ·
1 Parent(s): 54e984e

Refuse NaN/Inf at every int8 quantize; requant rounds half-up (retrained); bounded units; README for September

Browse files

Four local commits in one upload (ec432c8, c8db1cf, 47724cd, ae9cfa2): bounded elementwise units and LUT build (407 MB -> 2 MB, 277 MB -> 1 MB at startup); requant16 retrained from floor (-0.4981 LSB bias) to round-half-up, N/N 65,536/65,536, browser requant LUT re-exported; qat.py and every browser quantizer now refuse NaN/Inf instead of silently quantizing the whole layer to 0 (new tests fail on the old code); README September section. python test_verified_units.py 25/25, npm test 12/12 suites green.

README.md CHANGED
@@ -126,9 +126,10 @@ it is **not** a substitute for a real GPU on large models. Full envelope in
126
  differential gate.
127
  - **RDNA2 ISA audit hardenings** — bit-level (−0-aware) gate comparisons;
128
  proof that FMA contraction cannot change the quantize.
129
- - **Eleven-suite test chain** (`cd web && npm test`) — convergence, replicas,
130
  oracle, gates, properties, external corpus, **self-corpus** (the instruments
131
- scored against my own bugs), B2B, optimizer, transformer LM, int8 backward;
 
132
  results in [web/TEST_RESULTS.md](web/TEST_RESULTS.md).
133
  - **Dirty-buffer gate** — the pool is poisoned before a re-sweep so state bugs
134
  (a kernel assuming zeroed memory) are caught deterministically rather than
@@ -302,7 +303,7 @@ the gate; a stalled roster gradient is a bug in the *protocol*, where every
302
  computed value on every peer was correct, so no data oracle could fire. Those
303
  needed different instruments, not better oracles.
304
 
305
- All eleven test suites (`cd web && npm test`) pass; results with methodology in
306
  [`web/TEST_RESULTS.md`](web/TEST_RESULTS.md).
307
 
308
  ---
@@ -363,7 +364,7 @@ docker/ Dockerfile, dashboard image, compose (demo cluster)
363
  scripts/setup.bat / setup.sh interactive setup helpers
364
  config/ nodes + cluster env examples
365
  examples/my_task_template.py starting point for your own model
366
- test_verified_units.py full-domain checks of every shipped verified path (24)
367
  docs/ QUICKSTART, LIMITS, CUSTOM_TASK, TAILSCALE
368
  daisychain/spikewhale_task.py trains the real SpikeWhale on streamed HF datasets
369
  daisychain/spikewhale_panel.py slider control panel (localhost:8899)
@@ -371,6 +372,59 @@ web/ DaisyChain-Web: P2P browser training (WebRTC + WebG
371
  export_luts_web.py regenerates web/public LUTs from the trained units
372
  ```
373
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
374
  ## Recent updates (August 2026) — verified compute path
375
 
376
  Three changes to `daisychain/verified/`, each measured rather than asserted.
 
126
  differential gate.
127
  - **RDNA2 ISA audit hardenings** — bit-level (−0-aware) gate comparisons;
128
  proof that FMA contraction cannot change the quantize.
129
+ - **Twelve-suite test chain** (`cd web && npm test`) — convergence, replicas,
130
  oracle, gates, properties, external corpus, **self-corpus** (the instruments
131
+ scored against my own bugs), B2B, optimizer, transformer LM, int8 backward,
132
+ non-finite quantize refusal;
133
  results in [web/TEST_RESULTS.md](web/TEST_RESULTS.md).
134
  - **Dirty-buffer gate** — the pool is poisoned before a re-sweep so state bugs
135
  (a kernel assuming zeroed memory) are caught deterministically rather than
 
303
  computed value on every peer was correct, so no data oracle could fire. Those
304
  needed different instruments, not better oracles.
305
 
306
+ All twelve test suites (`cd web && npm test`) pass; results with methodology in
307
  [`web/TEST_RESULTS.md`](web/TEST_RESULTS.md).
308
 
309
  ---
 
364
  scripts/setup.bat / setup.sh interactive setup helpers
365
  config/ nodes + cluster env examples
366
  examples/my_task_template.py starting point for your own model
367
+ test_verified_units.py full-domain checks of every shipped verified path (25)
368
  docs/ QUICKSTART, LIMITS, CUSTOM_TASK, TAILSCALE
369
  daisychain/spikewhale_task.py trains the real SpikeWhale on streamed HF datasets
370
  daisychain/spikewhale_panel.py slider control panel (localhost:8899)
 
372
  export_luts_web.py regenerates web/public LUTs from the trained units
373
  ```
374
 
375
+ ## Recent updates (September 2026) — requant rounding, bounded units, non-finite guards
376
+
377
+ Everything below is measured, and `python test_verified_units.py` (25 checks) plus
378
+ `cd web && npm test` (twelve suites) reproduce it.
379
+
380
+ **The requant rounds instead of truncating.** `NeuralRequant16` was
381
+ `sat_int8(x >> 8)`, and Python's `>>` on negatives is an arithmetic floor, so every
382
+ requantized activation carried a systematic bias. Over the full int16 domain:
383
+
384
+ | requant | mean error | disagrees with the other on |
385
+ | --- | --- | --- |
386
+ | floor `x >> 8` (old) | **-0.4981 LSB** | 32,640 / 65,536 |
387
+ | round-half-up `(x + 128) >> 8` (now) | +0.0020 LSB | |
388
+
389
+ The shipped `requant16.pt` implemented the floor exactly, so this was real behaviour,
390
+ not a stale docstring. It was **retrained** against the new reference and saved only
391
+ at 65,536/65,536 (a 99.9% unit would silently break the LUT equivalence the browser
392
+ relies on), and `web/public/requant_lut.bin` was **re-exported** to match: the
393
+ invariant is *neural forward == LUT == native integer op*, so a table that lags its
394
+ unit is a fork. Tie rule, stated precisely: this is half-UP; Python/numpy/torch
395
+ `round` is half-to-even. The two differ on exactly the 128 exact ties (0.2% of
396
+ inputs) and both are unbiased to within 0.002 LSB. The Python fleet quantizes with
397
+ `np.round` and the browser with `floor(x + 0.5)`; the two fleets never co-train, so
398
+ that cannot fork a group.
399
+
400
+ **Bounded elementwise units.** The GEMM backends already capped their temporaries;
401
+ the units right after them did not. `relu_array` / `requant_array` received the
402
+ whole activation matrix in the proof path at ~1.5–2 KB per element (262,144
403
+ elements = **407 MB**), and the multiply-LUT build pushed 262,144 rows through the
404
+ atom net at startup (**277 MB RSS on every node**). Both are now blocked:
405
+ 407 MB → **2 MB, flat with N**; 277 MB → **1 MB**. Elementwise blocking is exact
406
+ (no accumulation order), verified bit-identical at N = 1 / 999 / 32768 / 32769 /
407
+ 100000. `dataset()` builders are vectorised (65,536 Python iterations → one numpy
408
+ pass, 0.015 s), which is what makes retraining a unit practical at all.
409
+
410
+ **A NaN or Inf no longer quantizes to zeros.** A float → int8 conversion has no
411
+ answer for NaN/Inf, and both trainers gave a silent wrong one:
412
+
413
+ | path | input | old result |
414
+ | --- | --- | --- |
415
+ | Python `qat.py` | `[0.5, nan, 3.0]` | scale = NaN → **`[0, 0, 0]`** (only a RuntimeWarning) |
416
+ | Python `qat.py` | `[0.5, inf, 3.0]` | scale = Inf → **`[0, 0, 0]`** |
417
+ | browser quantizers | `[0.5, Inf, 3]` | scale = Infinity → **`[0, 0, 0]`** |
418
+ | browser quantizers | `[0.5, NaN, 3]` | NaN skipped by the \|max\| scan, stored as **0** |
419
+
420
+ One bad value zeroed the whole layer's int8 input and training carried on. Now
421
+ `VerifiedLinear` raises `FloatingPointError`, and every browser quantizer
422
+ (`quantize`, `quantizeRows`, `quantizeCols`, `rowAbsMax`, `quantizeHeadCols`, and
423
+ the transformer's `quantizeColsAsRows`) throws `NonFiniteError`, naming how many
424
+ values were bad. Finite inputs are untouched, so builds with and without the guard
425
+ still co-train. Both new tests fail against the old code (Python: 2 checks; web:
426
+ `test_nonfinite.js`, 12) and pass now.
427
+
428
  ## Recent updates (August 2026) — verified compute path
429
 
430
  Three changes to `daisychain/verified/`, each measured rather than asserted.
daisychain/dashboard/agent.py CHANGED
@@ -19,6 +19,39 @@ STATUS_FILE = os.environ.get("DAISY_STATUS_FILE", "status.json")
19
  PORT = int(os.environ.get("DAISY_AGENT_PORT", "8900"))
20
 
21
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
22
  def _status():
23
  try:
24
  with open(STATUS_FILE) as f:
@@ -41,7 +74,7 @@ class Handler(BaseHTTPRequestHandler):
41
  if self.path.startswith("/health"):
42
  self._send({"ok": True, "rank": RANK, "host": socket.gethostname()})
43
  elif self.path.startswith("/resources"):
44
- r = survey_node(); r["rank"] = RANK; self._send(r)
45
  elif self.path.startswith("/status"):
46
  self._send(_status())
47
  else:
 
19
  PORT = int(os.environ.get("DAISY_AGENT_PORT", "8900"))
20
 
21
 
22
+ #: Cached capacity score. `survey_node()` defaults to measure=True, which runs
23
+ #: `capacity_score()` -- 0.3 s of timed 512x512 matmuls, measured at **0.47 s**
24
+ #: wall including warmup. The dashboard page carries
25
+ #: `<meta http-equiv="refresh" content="3">`, so every 3 s it re-scans every node
26
+ #: and each `/resources` request re-ran that benchmark: **~16% duty cycle of a
27
+ #: core, per node, continuously, while training** -- on hardware this project
28
+ #: targets precisely because it is small.
29
+ #:
30
+ #: It is also self-perturbing. The score feeds capacity-weighted batch sharding,
31
+ #: and was being measured while competing with both the training step and the
32
+ #: benchmark's own previous run, so the number it reported was a function of how
33
+ #: often the dashboard was open.
34
+ #:
35
+ #: Capacity is a property of the machine, not a live metric: measure once, reuse.
36
+ #: DAISY_CAPACITY still overrides (survey_node honours it), and a restart
37
+ #: re-measures.
38
+ _CAPACITY = None
39
+
40
+
41
+ def _resources():
42
+ """Node resources for /resources, without re-benchmarking on every request."""
43
+ global _CAPACITY
44
+ try:
45
+ if _CAPACITY is None:
46
+ _CAPACITY = survey_node().get("capacity") # measure once, at first ask
47
+ r = survey_node(measure=False) # cheap fields only
48
+ r["capacity"] = _CAPACITY
49
+ return r
50
+ except TypeError:
51
+ # the stdlib fallback survey_node above takes no `measure` kwarg
52
+ return survey_node()
53
+
54
+
55
  def _status():
56
  try:
57
  with open(STATUS_FILE) as f:
 
74
  if self.path.startswith("/health"):
75
  self._send({"ok": True, "rank": RANK, "host": socket.gethostname()})
76
  elif self.path.startswith("/resources"):
77
+ r = _resources(); r["rank"] = RANK; self._send(r)
78
  elif self.path.startswith("/status"):
79
  self._send(_status())
80
  else:
daisychain/dashboard/server.py CHANGED
@@ -2,6 +2,7 @@
2
  a readiness banner, P2P connectivity scan, pooled resource + capacity plan, and
3
  live training status. Config via DAISY_NODES_FILE. Serves on :8080."""
4
  import json
 
5
  import os
6
  from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
7
 
@@ -27,6 +28,18 @@ def _chip(ok, yes="OK", no="DOWN"):
27
  return f'<span class="px-2 py-0.5 rounded text-xs font-semibold {cls}">{yes if ok else no}</span>'
28
 
29
 
 
 
 
 
 
 
 
 
 
 
 
 
30
  def render(d):
31
  ready = d["ready"]
32
  banner = ("bg-emerald-500", "✓ CLUSTER READY — all nodes connected") if ready \
@@ -37,20 +50,20 @@ def render(d):
37
  res = n.get("resources") or {}
38
  dev = res.get("device", "—")
39
  gpu = (res.get("gpu") or {}).get("name", "")
40
- devlabel = f"{dev}" + (f" ({gpu})" if gpu else "")
41
  rows += f"""<tr class="border-b border-slate-100 dark:border-slate-800">
42
- <td class="py-2 px-3 font-medium">{n['name']}</td>
43
- <td class="py-2 px-3 text-slate-500">{n['host']}</td>
44
  <td class="py-2 px-3">{_chip(n['reachable'],'reachable','unreachable')}</td>
45
  <td class="py-2 px-3">{devlabel}</td>
46
  <td class="py-2 px-3 tabular-nums">{lat}</td></tr>"""
47
  plan = d["plan"]
48
  pr = ""
49
  for p in plan["per_node"]:
50
- dv = p["device"] + (f" ({p['gpu']})" if p.get("gpu") else "")
51
  ram = f'{p["ram_gb"]} GB' if p.get("ram_gb") is not None else "—"
52
  pr += f"""<tr class="border-b border-slate-100 dark:border-slate-800">
53
- <td class="py-2 px-3 font-medium">{p['name']}</td>
54
  <td class="py-2 px-3">{dv}</td>
55
  <td class="py-2 px-3 tabular-nums">{ram}</td>
56
  <td class="py-2 px-3 tabular-nums">{p['capacity']}</td>
 
2
  a readiness banner, P2P connectivity scan, pooled resource + capacity plan, and
3
  live training status. Config via DAISY_NODES_FILE. Serves on :8080."""
4
  import json
5
+ from html import escape as _html_escape
6
  import os
7
  from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
8
 
 
28
  return f'<span class="px-2 py-0.5 rounded text-xs font-semibold {cls}">{yes if ok else no}</span>'
29
 
30
 
31
+ #: Remote-supplied strings are escaped before they reach the page.
32
+ #:
33
+ #: `device` and the GPU `name` come from each node's /resources JSON -- i.e. from
34
+ #: the OTHER machine, not from local config -- and were interpolated raw into the
35
+ #: dashboard HTML. A node serving `{"gpu": {"name": "<script>...</script>"}}`
36
+ #: would run script in the operator's browser on the next 3-second refresh.
37
+ #: The node list itself is local config and trusted, but escaping it too costs
38
+ #: nothing and keeps the rule simple: nothing reaches the page unescaped.
39
+ def _e(v) -> str:
40
+ return _html_escape("" if v is None else str(v))
41
+
42
+
43
  def render(d):
44
  ready = d["ready"]
45
  banner = ("bg-emerald-500", "✓ CLUSTER READY — all nodes connected") if ready \
 
50
  res = n.get("resources") or {}
51
  dev = res.get("device", "—")
52
  gpu = (res.get("gpu") or {}).get("name", "")
53
+ devlabel = _e(dev) + (f" ({_e(gpu)})" if gpu else "")
54
  rows += f"""<tr class="border-b border-slate-100 dark:border-slate-800">
55
+ <td class="py-2 px-3 font-medium">{_e(n['name'])}</td>
56
+ <td class="py-2 px-3 text-slate-500">{_e(n['host'])}</td>
57
  <td class="py-2 px-3">{_chip(n['reachable'],'reachable','unreachable')}</td>
58
  <td class="py-2 px-3">{devlabel}</td>
59
  <td class="py-2 px-3 tabular-nums">{lat}</td></tr>"""
60
  plan = d["plan"]
61
  pr = ""
62
  for p in plan["per_node"]:
63
+ dv = _e(p["device"]) + (f" ({_e(p['gpu'])})" if p.get("gpu") else "")
64
  ram = f'{p["ram_gb"]} GB' if p.get("ram_gb") is not None else "—"
65
  pr += f"""<tr class="border-b border-slate-100 dark:border-slate-800">
66
+ <td class="py-2 px-3 font-medium">{_e(p['name'])}</td>
67
  <td class="py-2 px-3">{dv}</td>
68
  <td class="py-2 px-3 tabular-nums">{ram}</td>
69
  <td class="py-2 px-3 tabular-nums">{p['capacity']}</td>
daisychain/spikewhale_panel.py CHANGED
@@ -176,9 +176,24 @@ class H(BaseHTTPRequestHandler):
176
  pass
177
 
178
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
179
  def main():
180
- print(f"[spikewhale-panel] http://localhost:{PORT}", flush=True)
181
- ThreadingHTTPServer(("0.0.0.0", PORT), H).serve_forever()
 
182
 
183
 
184
  if __name__ == "__main__":
 
176
  pass
177
 
178
 
179
+ #: Bind address. LOOPBACK BY DEFAULT -- this is a control plane, not a status page.
180
+ #:
181
+ #: It bound 0.0.0.0 while printing "http://localhost:8899", so the operator was
182
+ #: told it was local-only when in fact anyone who could reach the port could
183
+ #: POST /start and spawn a training subprocess on this machine, unauthenticated.
184
+ #: `cfg["dataset"]` from that body becomes DAISY_SW_DATASET and is handed to
185
+ #: `datasets.load_dataset()`, so the caller also chooses what the box downloads.
186
+ #:
187
+ #: The agent and cluster dashboard legitimately listen on 0.0.0.0 -- they are
188
+ #: read-only telemetry that other nodes scan. This one starts and stops work, so
189
+ #: it defaults to loopback and exposure is opt-in via SW_PANEL_HOST.
190
+ HOST = os.environ.get("SW_PANEL_HOST", "127.0.0.1")
191
+
192
+
193
  def main():
194
+ print(f"[spikewhale-panel] http://{'localhost' if HOST in ('127.0.0.1', 'localhost') else HOST}:{PORT}"
195
+ + ("" if HOST == "127.0.0.1" else f" (listening on {HOST})"), flush=True)
196
+ ThreadingHTTPServer((HOST, PORT), H).serve_forever()
197
 
198
 
199
  if __name__ == "__main__":
daisychain/spikewhale_task.py CHANGED
@@ -24,7 +24,10 @@ import sys
24
 
25
  import torch
26
 
27
- _DEF_PATH = os.environ.get("DAISY_SW_PATH", r"C:\Users\quaz\Desktop\Spikewhale")
 
 
 
28
 
29
 
30
  def _import_spikewhale():
@@ -111,7 +114,6 @@ class SpikeWhaleTask:
111
  rank = _envi("RANK", 0); world = _envi("WORLD_SIZE", 1)
112
  self.stream = _FineWebStream(self.tok, self.seqlen, rank, world)
113
  self._SpikeWhaleLM = SpikeWhaleLM
114
- n = None
115
  print(f"[spikewhale] data source: {self.stream.source}", flush=True)
116
 
117
  def build_model(self):
 
24
 
25
  import torch
26
 
27
+ #: Where model_v2.py / config.py / tokenizer.json live. The default used to be an
28
+ #: absolute path on the author's own desktop, which in a published repo means every
29
+ #: other user gets an ImportError naming a directory that was never theirs.
30
+ _DEF_PATH = os.environ.get("DAISY_SW_PATH", os.path.join(os.getcwd(), "Spikewhale"))
31
 
32
 
33
  def _import_spikewhale():
 
114
  rank = _envi("RANK", 0); world = _envi("WORLD_SIZE", 1)
115
  self.stream = _FineWebStream(self.tok, self.seqlen, rank, world)
116
  self._SpikeWhaleLM = SpikeWhaleLM
 
117
  print(f"[spikewhale] data source: {self.stream.source}", flush=True)
118
 
119
  def build_model(self):
daisychain/verified/common.py CHANGED
@@ -6,6 +6,7 @@ entire finite input domain* (N/N verification).
6
  """
7
  from __future__ import annotations
8
 
 
9
  import torch
10
  import torch.nn as nn
11
 
@@ -29,6 +30,21 @@ def pm(bits: torch.Tensor) -> torch.Tensor:
29
  return bits * 2.0 - 1.0
30
 
31
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
32
  def mlp(inp: int, out: int, h: int = 128, layers: int = 3) -> nn.Sequential:
33
  mods: list[nn.Module] = [nn.Linear(inp, h), nn.ReLU()]
34
  for _ in range(layers - 1):
 
6
  """
7
  from __future__ import annotations
8
 
9
+ import numpy as np
10
  import torch
11
  import torch.nn as nn
12
 
 
30
  return bits * 2.0 - 1.0
31
 
32
 
33
+ def bits_matrix(vals, n: int) -> torch.Tensor:
34
+ """LSB-first bit rows for a whole array of ints -> float32 (len(vals), n).
35
+
36
+ Vectorised counterpart of `bits_of`. The per-unit `dataset()` builders used a
37
+ Python loop with two `bits_of` calls and a `torch.stack` per entry; for
38
+ NeuralRequant16 that is 65,536 iterations building 131,072 tiny tensors, and
39
+ it dominates both `fit()` and `verify()` -- the two things you must run to
40
+ change a unit at all. numpy builds the same matrix in one pass.
41
+
42
+ Identical output to stacking `bits_of` row by row.
43
+ """
44
+ a = np.asarray(vals, dtype=np.int64).reshape(-1, 1)
45
+ return torch.from_numpy(((a >> np.arange(n)) & 1).astype(np.float32))
46
+
47
+
48
  def mlp(inp: int, out: int, h: int = 128, layers: int = 3) -> nn.Sequential:
49
  mods: list[nn.Module] = [nn.Linear(inp, h), nn.ReLU()]
50
  for _ in range(layers - 1):
daisychain/verified/mul8.py CHANGED
@@ -58,14 +58,23 @@ class NeuralMul4:
58
  a = (np.asarray(a).astype(np.int64) & 0xF)
59
  b = (np.asarray(b).astype(np.int64) & 0xF)
60
  idx = np.arange(4)
61
- bits_a = (a[:, None] >> idx) & 1
62
- bits_b = (b[:, None] >> idx) & 1
63
- x = np.concatenate([bits_a, bits_b], axis=1).astype(np.float32) * 2.0 - 1.0
64
- out = (self.net(torch.from_numpy(x)) > 0).to(torch.int64).numpy()
65
  from . import instrument
66
- instrument.bump("NeuralMul4.forward_calls", 1)
67
- instrument.bump("NeuralMul4.products", a.shape[0])
68
- return (out * (1 << np.arange(8))).sum(axis=1)
 
 
 
 
 
 
 
 
 
 
 
 
 
69
 
70
 
71
  class NeuralMul8:
 
58
  a = (np.asarray(a).astype(np.int64) & 0xF)
59
  b = (np.asarray(b).astype(np.int64) & 0xF)
60
  idx = np.arange(4)
 
 
 
 
61
  from . import instrument
62
+ from .ops import _blocked
63
+
64
+ def _run(sl):
65
+ ca, cb = a[sl], b[sl]
66
+ bits_a = (ca[:, None] >> idx) & 1
67
+ bits_b = (cb[:, None] >> idx) & 1
68
+ x = np.concatenate([bits_a, bits_b], axis=1).astype(np.float32) * 2.0 - 1.0
69
+ out = (self.net(torch.from_numpy(x)) > 0).to(torch.int64).numpy()
70
+ instrument.bump("NeuralMul4.forward_calls", 1)
71
+ return (out * (1 << np.arange(8))).sum(axis=1)
72
+
73
+ # BOUNDED. `NeuralBackend.gemm` blocks before it reaches here, but the LUT build does not:
74
+ # `build_mul8_lut` hands NeuralMul8.mul_array all 65,536 pairs, which become 262,144 rows
75
+ # through this net. Measured 277 MB RSS at import time, on every node, on hardware this
76
+ # project targets for being small. Blocking is exact -- this is an elementwise map.
77
+ return _blocked(_run, np.arange(a.shape[0]))
78
 
79
 
80
  class NeuralMul8:
daisychain/verified/ops.py CHANGED
@@ -17,10 +17,44 @@ from __future__ import annotations
17
  import numpy as np
18
  import torch
19
 
20
- from .common import bits_of, int_of, pm, mlp, verify, train
21
  from . import instrument
22
 
23
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
24
  def _s(v: int, bits: int) -> int:
25
  """two's-complement raw -> signed."""
26
  m = 1 << (bits - 1)
@@ -38,11 +72,9 @@ class NeuralReLU8:
38
  self.net = mlp(8, 8, h=h, layers=layers)
39
 
40
  def dataset(self):
41
- X, Y = [], []
42
- for b in range(256):
43
- X.append(pm(bits_of(b, 8)))
44
- Y.append(bits_of(max(0, _s(b, 8)) & 0xFF, 8))
45
- return torch.stack(X), torch.stack(Y)
46
 
47
  def fit(self, steps: int = 3000, lr: float = 2e-3, tag: str = "relu8"):
48
  X, Y = self.dataset(); train(self.net, X, Y, steps=steps, lr=lr, tag=tag); return self
@@ -61,10 +93,14 @@ class NeuralReLU8:
61
  """Batched int8 ReLU over an array (one neural forward)."""
62
  self.net.eval()
63
  a = np.asarray(arr).astype(np.int64).ravel() & 0xFF
64
- bits = ((a[:, None] >> np.arange(8)) & 1).astype(np.float32) * 2.0 - 1.0
65
- out = (self.net(torch.from_numpy(bits)) > 0).to(torch.int64).numpy()
66
- raw = (out * (1 << np.arange(8))).sum(axis=1)
67
- instrument.bump("NeuralReLU8.forward_calls", 1)
 
 
 
 
68
  instrument.bump("NeuralReLU8.elements", a.shape[0])
69
  return np.where(raw >= 128, raw - 256, raw).reshape(np.asarray(arr).shape)
70
 
@@ -77,14 +113,36 @@ class NeuralRequant16:
77
  self.net = mlp(16, 8, h=h, layers=layers)
78
 
79
  def _ref(self, x_signed: int) -> int:
80
- return sat_int8(x_signed >> self.shift) # arithmetic shift, saturate
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
81
 
82
  def dataset(self):
83
- X, Y = [], []
84
- for b in range(65536):
85
- X.append(pm(bits_of(b, 16)))
86
- Y.append(bits_of(self._ref(_s(b, 16)) & 0xFF, 8))
87
- return torch.stack(X), torch.stack(Y)
88
 
89
  def fit(self, steps: int = 6000, lr: float = 2e-3, tag: str = "requant16"):
90
  X, Y = self.dataset(); train(self.net, X, Y, steps=steps, lr=lr, tag=tag); return self
@@ -103,9 +161,13 @@ class NeuralRequant16:
103
  """Batched int16->int8 requantize over an array (one neural forward)."""
104
  self.net.eval()
105
  a = np.asarray(arr).astype(np.int64).ravel() & 0xFFFF
106
- bits = ((a[:, None] >> np.arange(16)) & 1).astype(np.float32) * 2.0 - 1.0
107
- out = (self.net(torch.from_numpy(bits)) > 0).to(torch.int64).numpy()
108
- raw = (out * (1 << np.arange(8))).sum(axis=1)
109
- instrument.bump("NeuralRequant16.forward_calls", 1)
 
 
 
 
110
  instrument.bump("NeuralRequant16.elements", a.shape[0])
111
  return np.where(raw >= 128, raw - 256, raw).reshape(np.asarray(arr).shape)
 
17
  import numpy as np
18
  import torch
19
 
20
+ from .common import bits_of, bits_matrix, int_of, pm, mlp, verify, train
21
  from . import instrument
22
 
23
 
24
+ #: Cap on ELEMENTS handed to one batched neural forward, for the two elementwise
25
+ #: units below.
26
+ #:
27
+ #: `NeuralBackend.gemm` and `LUTBackend.gemm` both bound their temporaries; these
28
+ #: two did not, and they sit directly downstream of the GEMM in the proof path
29
+ #: (`_VerifiedQGEMM.forward` -> `requant_array(acc16)` -> `relu_array(yq)`), where
30
+ #: they receive the WHOLE activation matrix at once. `requant_array` widens each
31
+ #: element to 16 float32 inputs and pushes them through a 3-layer, 256-wide MLP,
32
+ #: so the hidden activations alone are N*256*4 bytes per layer.
33
+ #:
34
+ #: Measured (RSS delta, requant16): ~1.5-2 KB per element -- 65,536 elements cost
35
+ #: 144 MB and 262,144 cost 407 MB. A 512x768 activation is 393,216 elements. That
36
+ #: is the same failure the GEMM caps already guard against, on hardware this
37
+ #: project targets precisely because it is small.
38
+ #:
39
+ #: 32,768 elements ~= 64 MB, comparable to NeuralBackend's own product budget.
40
+ #:
41
+ #: Blocking is UNCONDITIONALLY exact here in a way it is not for a GEMM: these are
42
+ #: elementwise maps, so no accumulation order exists to perturb. Every output byte
43
+ #: is identical to the unblocked path.
44
+ MAX_UNIT_ELEMENTS = 1 << 15
45
+
46
+
47
+ def _blocked(fn, a: np.ndarray, cap: int = MAX_UNIT_ELEMENTS) -> np.ndarray:
48
+ """Apply an elementwise unit to a flat array in bounded chunks."""
49
+ n = a.shape[0]
50
+ if n <= cap:
51
+ return fn(a)
52
+ out = np.empty(n, dtype=np.int64)
53
+ for i in range(0, n, cap):
54
+ out[i:i + cap] = fn(a[i:i + cap])
55
+ return out
56
+
57
+
58
  def _s(v: int, bits: int) -> int:
59
  """two's-complement raw -> signed."""
60
  m = 1 << (bits - 1)
 
72
  self.net = mlp(8, 8, h=h, layers=layers)
73
 
74
  def dataset(self):
75
+ b = np.arange(256)
76
+ sgn = np.where(b >= 128, b - 256, b)
77
+ return pm(bits_matrix(b, 8)), bits_matrix(np.maximum(0, sgn) & 0xFF, 8)
 
 
78
 
79
  def fit(self, steps: int = 3000, lr: float = 2e-3, tag: str = "relu8"):
80
  X, Y = self.dataset(); train(self.net, X, Y, steps=steps, lr=lr, tag=tag); return self
 
93
  """Batched int8 ReLU over an array (one neural forward)."""
94
  self.net.eval()
95
  a = np.asarray(arr).astype(np.int64).ravel() & 0xFF
96
+
97
+ def _run(chunk):
98
+ bits = ((chunk[:, None] >> np.arange(8)) & 1).astype(np.float32) * 2.0 - 1.0
99
+ out = (self.net(torch.from_numpy(bits)) > 0).to(torch.int64).numpy()
100
+ instrument.bump("NeuralReLU8.forward_calls", 1)
101
+ return (out * (1 << np.arange(8))).sum(axis=1)
102
+
103
+ raw = _blocked(_run, a) # bounded; see MAX_UNIT_ELEMENTS
104
  instrument.bump("NeuralReLU8.elements", a.shape[0])
105
  return np.where(raw >= 128, raw - 256, raw).reshape(np.asarray(arr).shape)
106
 
 
113
  self.net = mlp(16, 8, h=h, layers=layers)
114
 
115
  def _ref(self, x_signed: int) -> int:
116
+ """sat_int8(round_half_up(x / 2**shift)) -- rounding, not truncation.
117
+
118
+ This was `x >> shift`, and Python's `>>` on negatives is an arithmetic
119
+ FLOOR. Measured over the full int16 domain (unsaturated band): floor has a
120
+ mean error of **-0.4981 LSB** -- a systematic negative bias on every
121
+ requantized activation in the network -- against +0.002 for round-half-up.
122
+ The two disagree on 32,640 of 65,536 inputs.
123
+
124
+ Half-up via `(x + 2**(shift-1)) >> shift` is exact integer arithmetic and
125
+ is the canonical hardware requant, which matters because this unit exists
126
+ to model hardware. It removes the floor bias the Byrne/morpho work avoids with
127
+ `sat_int8(round(acc / 128))` -- but note the TIE rule: Python/numpy/torch
128
+ `round` is half-to-EVEN, this is half-UP. Measured over the int16 domain they
129
+ disagree on exactly the 128 ties (x mod 256 == 128 with an even quotient),
130
+ 0.2% of inputs; both are unbiased to within 0.002 LSB. So "agree" holds for
131
+ the bias, not bit-for-bit unless morpho also rounds half-up. (Same tie
132
+ question: `qat.py` quantizes with np.round, half-even; the browser build uses
133
+ floor(x + 0.5), half-up. The two fleets never co-train, so they cannot fork.)
134
+
135
+ Changing this invalidates any previously trained requant16.pt: the unit is
136
+ verified N/N against THIS reference, so the weights must be retrained and
137
+ re-verified whenever it moves.
138
+ """
139
+ return sat_int8((x_signed + (1 << (self.shift - 1))) >> self.shift)
140
 
141
  def dataset(self):
142
+ b = np.arange(65536)
143
+ sgn = np.where(b >= 32768, b - 65536, b)
144
+ y = np.clip((sgn + (1 << (self.shift - 1))) >> self.shift, -128, 127)
145
+ return pm(bits_matrix(b, 16)), bits_matrix(y & 0xFF, 8)
 
146
 
147
  def fit(self, steps: int = 6000, lr: float = 2e-3, tag: str = "requant16"):
148
  X, Y = self.dataset(); train(self.net, X, Y, steps=steps, lr=lr, tag=tag); return self
 
161
  """Batched int16->int8 requantize over an array (one neural forward)."""
162
  self.net.eval()
163
  a = np.asarray(arr).astype(np.int64).ravel() & 0xFFFF
164
+
165
+ def _run(chunk):
166
+ bits = ((chunk[:, None] >> np.arange(16)) & 1).astype(np.float32) * 2.0 - 1.0
167
+ out = (self.net(torch.from_numpy(bits)) > 0).to(torch.int64).numpy()
168
+ instrument.bump("NeuralRequant16.forward_calls", 1)
169
+ return (out * (1 << np.arange(8))).sum(axis=1)
170
+
171
+ raw = _blocked(_run, a) # bounded; see MAX_UNIT_ELEMENTS
172
  instrument.bump("NeuralRequant16.elements", a.shape[0])
173
  return np.where(raw >= 128, raw - 256, raw).reshape(np.asarray(arr).shape)
daisychain/verified/qat.py CHANGED
@@ -1,116 +1,126 @@
1
- """Quantization-aware training THROUGH the verified units.
2
-
3
- This is the piece that was missing: a trainable layer whose FORWARD compute
4
- actually runs on the verified GUDA logic --
5
-
6
- quantize -> NeuralMul (verified INT8 multiply) GEMM -> NeuralRequant16 ->
7
- NeuralReLU8 -> dequantize
8
-
9
- -- while the BACKWARD uses a straight-through estimator (the integer path has no
10
- gradient), so ordinary float weights still learn. With instrument.enable(), each
11
- unit records how many times its neural forward ran, so a training run leaves
12
- hard evidence (call counts) that it computed through the units, not around them.
13
-
14
- Honest cost: every forward multiply is a neural forward pass -> this is SLOW
15
- (functional, not fast). It is a correctness/《evidence》demo, not a speed path.
16
- """
17
- from __future__ import annotations
18
-
19
- import numpy as np
20
- import torch
21
- import torch.nn as nn
22
-
23
- from .backends import NeuralBackend
24
-
25
-
26
- class _VerifiedQGEMM(torch.autograd.Function):
27
- @staticmethod
28
- def forward(ctx, x, w, mul, requant, relu_unit, use_relu, luts):
29
- ctx.save_for_backward(x, w)
30
- ctx.device = x.device # verified units run on CPU;
31
- xnp, wnp = x.detach().cpu().numpy(), w.detach().cpu().numpy()
32
- sx = max(float(np.abs(xnp).max()) / 127.0, 1e-8)
33
- sw = max(float(np.abs(wnp).max()) / 127.0, 1e-8)
34
- xq = np.clip(np.round(xnp / sx), -128, 127).astype(np.int8)
35
- wq = np.clip(np.round(wnp / sw), -128, 127).astype(np.int8)
36
-
37
- if luts is not None:
38
- # FAST path: the verified units, materialized as lookup tables
39
- # (bit-identical to the neural forward, ~500x faster).
40
- from . import instrument
41
- acc = luts["backend"].gemm(xq, wq) # counts LUT products
42
- acc16 = np.clip(acc, -32768, 32767).astype(np.int64)
43
- yq = luts["requant"][acc16 & 0xFFFF]
44
- instrument.bump("VerifiedRequant16(LUT).elements", acc16.size)
45
- if use_relu:
46
- yq = luts["relu"][yq & 0xFF]
47
- instrument.bump("VerifiedReLU8(LUT).elements", yq.size)
48
- else:
49
- acc = NeuralBackend(mul).gemm(xq, wq) # verified multiply fires
50
- acc16 = np.clip(acc, -32768, 32767).astype(np.int64)
51
- yq = requant.requant_array(acc16) # requant16 fires
52
- if use_relu:
53
- yq = relu_unit.relu_array(yq) # relu8 fires
54
- dequant = sx * sw * 256.0 # undo requant's >>8
55
- return torch.from_numpy(yq.astype(np.float32) * dequant).to(ctx.device)
56
-
57
- @staticmethod
58
- def backward(ctx, gy):
59
- # straight-through: treat the quantized path as y ≈ x @ w
60
- x, w = ctx.saved_tensors
61
- return gy @ w.t(), x.t() @ gy, None, None, None, None, None
62
-
63
-
64
- def build_luts(mul, requant, relu_unit):
65
- """Materialize the verified units as lookup tables (one-time). The result is
66
- bit-identical to the neural forward but ~500x faster to run."""
67
- from .lut import LUTBackend, build_requant16_lut, build_relu8_lut
68
- return {"backend": LUTBackend(mul),
69
- "requant": build_requant16_lut(requant),
70
- "relu": build_relu8_lut(relu_unit)}
71
-
72
-
73
- class VerifiedLinear(nn.Module):
74
- """Linear layer whose forward is computed by the verified units.
75
-
76
- fast=True materializes the units as LUTs (bit-identical, ~500x faster) so
77
- verified training is practical; fast=False runs the neural forward (proof).
78
- """
79
-
80
- def __init__(self, in_f, out_f, mul, requant, relu_unit, use_relu=True,
81
- fast=False):
82
- super().__init__()
83
- self.weight = nn.Parameter(torch.randn(in_f, out_f) * 0.3)
84
- self.bias = nn.Parameter(torch.zeros(out_f))
85
- self.mul, self.requant, self.relu_unit = mul, requant, relu_unit
86
- self.use_relu = use_relu
87
- self.luts = build_luts(mul, requant, relu_unit) if fast else None
88
-
89
- def forward(self, x):
90
- y = _VerifiedQGEMM.apply(x, self.weight, self.mul, self.requant,
91
- self.relu_unit, self.use_relu, self.luts)
92
- return y + self.bias
93
-
94
-
95
- def _weights_dir():
96
- import os
97
- return os.path.join(os.path.dirname(os.path.abspath(__file__)), "weights")
98
-
99
-
100
- def load_units(mul_pt=None, requant_pt=None, relu_pt=None):
101
- """Load the TRAINED, N/N-verified units bundled with DaisyChain."""
102
- import os
103
- wd = _weights_dir()
104
- mul_pt = mul_pt or os.path.join(wd, "mul8.pt")
105
- requant_pt = requant_pt or os.path.join(wd, "requant16.pt")
106
- relu_pt = relu_pt or os.path.join(wd, "relu8.pt")
107
- from .mul8 import NeuralMul8
108
- from .ops import NeuralReLU8, NeuralRequant16
109
- mul = NeuralMul8()
110
- mul.atom.net.load_state_dict(torch.load(mul_pt)["state_dict"]); mul.atom.net.eval()
111
- relu = NeuralReLU8()
112
- relu.net.load_state_dict(torch.load(relu_pt)["state_dict"]); relu.net.eval()
113
- ck = torch.load(requant_pt)
114
- rq = NeuralRequant16(shift=ck["shift"])
115
- rq.net.load_state_dict(ck["state_dict"]); rq.net.eval()
116
- return mul, rq, relu
 
 
 
 
 
 
 
 
 
 
 
1
+ """Quantization-aware training THROUGH the verified units.
2
+
3
+ This is the piece that was missing: a trainable layer whose FORWARD compute
4
+ actually runs on the verified GUDA logic --
5
+
6
+ quantize -> NeuralMul (verified INT8 multiply) GEMM -> NeuralRequant16 ->
7
+ NeuralReLU8 -> dequantize
8
+
9
+ -- while the BACKWARD uses a straight-through estimator (the integer path has no
10
+ gradient), so ordinary float weights still learn. With instrument.enable(), each
11
+ unit records how many times its neural forward ran, so a training run leaves
12
+ hard evidence (call counts) that it computed through the units, not around them.
13
+
14
+ Honest cost: every forward multiply is a neural forward pass -> this is SLOW
15
+ (functional, not fast). It is a correctness/《evidence》demo, not a speed path.
16
+ """
17
+ from __future__ import annotations
18
+
19
+ import numpy as np
20
+ import torch
21
+ import torch.nn as nn
22
+
23
+ from .backends import NeuralBackend
24
+
25
+
26
+ class _VerifiedQGEMM(torch.autograd.Function):
27
+ @staticmethod
28
+ def forward(ctx, x, w, mul, requant, relu_unit, use_relu, luts):
29
+ ctx.save_for_backward(x, w)
30
+ ctx.device = x.device # verified units run on CPU;
31
+ xnp, wnp = x.detach().cpu().numpy(), w.detach().cpu().numpy()
32
+ # A float -> int8 cast of NaN/Inf is undefined; numpy returns 0 with only a
33
+ # RuntimeWarning. One non-finite value anywhere made the whole-tensor scale NaN/Inf
34
+ # and silently zeroed EVERY quantized input of the layer (measured: [0.5, nan, 3.0]
35
+ # -> [0, 0, 0]). Refuse at the source so a diverging run fails where it diverged.
36
+ for nm, a in (("activations", xnp), ("weights", wnp)):
37
+ if not np.isfinite(a).all():
38
+ raise FloatingPointError(
39
+ "non-finite %s entering the verified GEMM (%d of %d values); the int8 "
40
+ "quantize would silently turn them all into 0"
41
+ % (nm, int((~np.isfinite(a)).sum()), a.size))
42
+ sx = max(float(np.abs(xnp).max()) / 127.0, 1e-8)
43
+ sw = max(float(np.abs(wnp).max()) / 127.0, 1e-8)
44
+ xq = np.clip(np.round(xnp / sx), -128, 127).astype(np.int8)
45
+ wq = np.clip(np.round(wnp / sw), -128, 127).astype(np.int8)
46
+
47
+ if luts is not None:
48
+ # FAST path: the verified units, materialized as lookup tables
49
+ # (bit-identical to the neural forward, ~500x faster).
50
+ from . import instrument
51
+ acc = luts["backend"].gemm(xq, wq) # counts LUT products
52
+ acc16 = np.clip(acc, -32768, 32767).astype(np.int64)
53
+ yq = luts["requant"][acc16 & 0xFFFF]
54
+ instrument.bump("VerifiedRequant16(LUT).elements", acc16.size)
55
+ if use_relu:
56
+ yq = luts["relu"][yq & 0xFF]
57
+ instrument.bump("VerifiedReLU8(LUT).elements", yq.size)
58
+ else:
59
+ acc = NeuralBackend(mul).gemm(xq, wq) # verified multiply fires
60
+ acc16 = np.clip(acc, -32768, 32767).astype(np.int64)
61
+ yq = requant.requant_array(acc16) # requant16 fires
62
+ if use_relu:
63
+ yq = relu_unit.relu_array(yq) # relu8 fires
64
+ dequant = sx * sw * 256.0 # undo requant's >>8
65
+ return torch.from_numpy(yq.astype(np.float32) * dequant).to(ctx.device)
66
+
67
+ @staticmethod
68
+ def backward(ctx, gy):
69
+ # straight-through: treat the quantized path as y ≈ x @ w
70
+ x, w = ctx.saved_tensors
71
+ return gy @ w.t(), x.t() @ gy, None, None, None, None, None
72
+
73
+
74
+ def build_luts(mul, requant, relu_unit):
75
+ """Materialize the verified units as lookup tables (one-time). The result is
76
+ bit-identical to the neural forward but ~500x faster to run."""
77
+ from .lut import LUTBackend, build_requant16_lut, build_relu8_lut
78
+ return {"backend": LUTBackend(mul),
79
+ "requant": build_requant16_lut(requant),
80
+ "relu": build_relu8_lut(relu_unit)}
81
+
82
+
83
+ class VerifiedLinear(nn.Module):
84
+ """Linear layer whose forward is computed by the verified units.
85
+
86
+ fast=True materializes the units as LUTs (bit-identical, ~500x faster) so
87
+ verified training is practical; fast=False runs the neural forward (proof).
88
+ """
89
+
90
+ def __init__(self, in_f, out_f, mul, requant, relu_unit, use_relu=True,
91
+ fast=False):
92
+ super().__init__()
93
+ self.weight = nn.Parameter(torch.randn(in_f, out_f) * 0.3)
94
+ self.bias = nn.Parameter(torch.zeros(out_f))
95
+ self.mul, self.requant, self.relu_unit = mul, requant, relu_unit
96
+ self.use_relu = use_relu
97
+ self.luts = build_luts(mul, requant, relu_unit) if fast else None
98
+
99
+ def forward(self, x):
100
+ y = _VerifiedQGEMM.apply(x, self.weight, self.mul, self.requant,
101
+ self.relu_unit, self.use_relu, self.luts)
102
+ return y + self.bias
103
+
104
+
105
+ def _weights_dir():
106
+ import os
107
+ return os.path.join(os.path.dirname(os.path.abspath(__file__)), "weights")
108
+
109
+
110
+ def load_units(mul_pt=None, requant_pt=None, relu_pt=None):
111
+ """Load the TRAINED, N/N-verified units bundled with DaisyChain."""
112
+ import os
113
+ wd = _weights_dir()
114
+ mul_pt = mul_pt or os.path.join(wd, "mul8.pt")
115
+ requant_pt = requant_pt or os.path.join(wd, "requant16.pt")
116
+ relu_pt = relu_pt or os.path.join(wd, "relu8.pt")
117
+ from .mul8 import NeuralMul8
118
+ from .ops import NeuralReLU8, NeuralRequant16
119
+ mul = NeuralMul8()
120
+ mul.atom.net.load_state_dict(torch.load(mul_pt)["state_dict"]); mul.atom.net.eval()
121
+ relu = NeuralReLU8()
122
+ relu.net.load_state_dict(torch.load(relu_pt)["state_dict"]); relu.net.eval()
123
+ ck = torch.load(requant_pt)
124
+ rq = NeuralRequant16(shift=ck["shift"])
125
+ rq.net.load_state_dict(ck["state_dict"]); rq.net.eval()
126
+ return mul, rq, relu
daisychain/verified/weights/requant16.pt CHANGED
@@ -1,3 +1,3 @@
1
  version https://git-lfs.github.com/spec/v1
2
- oid sha256:0cca8f64e40b4969395b588d781a4ae49bca0ede8d13576ea03dbe0fe1cb2ea2
3
  size 555160
 
1
  version https://git-lfs.github.com/spec/v1
2
+ oid sha256:1ea8a556f59a94b5d155ccd8b9476c0976f29b578f2cecfeff269a1a3292f750
3
  size 555160
export_luts_web.py CHANGED
@@ -32,12 +32,37 @@ def main():
32
  with open(os.path.join(OUT, "luts_meta.json"), "w") as f:
33
  json.dump(meta, f)
34
 
35
- # sanity: LUT must equal the true signed product (verified units are exact)
36
- a, b = 37, -19
37
- au, bu = a & 0xFF, b & 0xFF
38
- assert int(mul_lut[au, bu]) == a * b, "mul LUT mismatch"
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
39
  print("exported mul_lut(int16 65536), requant_lut(int8 65536), relu_lut(int8 256)")
40
- print("shift =", rq.shift, "| sanity 37*-19 =", int(mul_lut[au, bu]))
 
41
 
42
 
43
  if __name__ == "__main__":
 
32
  with open(os.path.join(OUT, "luts_meta.json"), "w") as f:
33
  json.dump(meta, f)
34
 
35
+ # CERTIFY THE EXPORTED TABLES, not a sample of them.
36
+ #
37
+ # These binaries are what the BROWSER computes through, and each is written
38
+ # after a narrowing cast (int64 -> int16 / int8) that numpy performs silently:
39
+ # an out-of-range value wraps rather than raising. The previous check was a
40
+ # single pair (37 * -19), which cannot see a wrap anywhere else in the domain.
41
+ #
42
+ # The domain is finite and tiny, so checking ALL of it is the exhaustive
43
+ # verification rather than a sample -- ~0.5 ms, the same argument
44
+ # `certify_mul8_lut` already makes for the runtime table.
45
+ au = np.repeat(np.arange(256), 256)
46
+ bu = np.tile(np.arange(256), 256)
47
+ sa = np.where(au >= 128, au - 256, au)
48
+ sb = np.where(bu >= 128, bu - 256, bu)
49
+ bad = int((mul_lut[au, bu].astype(np.int64) != sa * sb).sum())
50
+ if bad:
51
+ raise SystemExit("mul LUT: %d/65536 entries wrong after int16 cast" % bad)
52
+
53
+ ref_rq = build_requant16_lut(rq).astype(np.int64)
54
+ bad_rq = int((req_lut.astype(np.int64) != ref_rq).sum())
55
+ if bad_rq:
56
+ raise SystemExit("requant LUT: %d/65536 entries wrong after int8 cast" % bad_rq)
57
+
58
+ ref_relu = build_relu8_lut(relu).astype(np.int64)
59
+ bad_relu = int((relu_lut.astype(np.int64) != ref_relu).sum())
60
+ if bad_relu:
61
+ raise SystemExit("relu LUT: %d/256 entries wrong after int8 cast" % bad_relu)
62
+
63
  print("exported mul_lut(int16 65536), requant_lut(int8 65536), relu_lut(int8 256)")
64
+ print("certified after cast: mul 65536/65536, requant 65536/65536, relu 256/256")
65
+ print("shift =", rq.shift)
66
 
67
 
68
  if __name__ == "__main__":
test_verified_units.py CHANGED
@@ -58,11 +58,16 @@ def main():
58
  ck("LUTBackend self-certified at construction",
59
  getattr(backend, "certified", None) == (65536, 65536))
60
 
61
- print("requantize -- sat_int8(x >> 8), full int16 domain")
62
  x = np.arange(65536)
63
  xs = np.where(x >= 32768, x - 65536, x)
64
- rq_gold = np.clip(xs >> requant.shift, -128, 127)
65
- ck("requant_array == sat_int8(x >> shift) over all 65536",
 
 
 
 
 
66
  np.array_equal(requant.requant_array(x), rq_gold))
67
  ck("requant LUT == golden over all 65536",
68
  np.array_equal(luts["requant"][x & 0xFFFF], rq_gold))
@@ -155,6 +160,20 @@ def main():
155
  ck("require() fails when a unit is under-invoked", True)
156
  instrument.disable()
157
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
158
  print()
159
  if FAILURES:
160
  print("FAILED: %d" % len(FAILURES))
 
58
  ck("LUTBackend self-certified at construction",
59
  getattr(backend, "certified", None) == (65536, 65536))
60
 
61
+ print("requantize -- sat_int8(round_half_up(x / 256)), full int16 domain")
62
  x = np.arange(65536)
63
  xs = np.where(x >= 32768, x - 65536, x)
64
+ # Golden is ROUND-HALF-UP, not the arithmetic shift this used to assert.
65
+ # `x >> shift` is a floor, which carries a systematic -0.4981 LSB bias on every
66
+ # requantized activation (measured over this same domain); half-up measures
67
+ # +0.0020. The unit's reference was corrected and requant16.pt retrained
68
+ # against it, so this golden moved with it -- see NeuralRequant16._ref.
69
+ rq_gold = np.clip((xs + (1 << (requant.shift - 1))) >> requant.shift, -128, 127)
70
+ ck("requant_array == sat_int8(round_half_up(x/2**shift)) over all 65536",
71
  np.array_equal(requant.requant_array(x), rq_gold))
72
  ck("requant LUT == golden over all 65536",
73
  np.array_equal(luts["requant"][x & 0xFFFF], rq_gold))
 
160
  ck("require() fails when a unit is under-invoked", True)
161
  instrument.disable()
162
 
163
+ print("quantize -- a non-finite input must fail loudly, not become int8 zeros")
164
+ import torch
165
+ from daisychain.verified.qat import VerifiedLinear
166
+ layer = VerifiedLinear(3, 2, mul, requant, relu_unit, fast=True)
167
+ ok_out = layer(torch.tensor([[0.5, -1.0, 3.0]]))
168
+ ck("a finite batch still runs", bool(torch.isfinite(ok_out).all()))
169
+ for bad in (float("nan"), float("inf")):
170
+ try:
171
+ layer(torch.tensor([[0.5, bad, 3.0]]))
172
+ ck("%r in the activations raises" % bad, False,
173
+ "quantized silently (the old path returned all-zero int8)")
174
+ except FloatingPointError:
175
+ ck("%r in the activations raises" % bad, True)
176
+
177
  print()
178
  if FAILURES:
179
  print("FAILED: %d" % len(FAILURES))
web/TEST_RESULTS.md CHANGED
@@ -1,6 +1,6 @@
1
  # Test results — updated 2026-07-17
2
 
3
- All suites run with `npm test` (chains all eleven). Every suite exits 0.
4
  Hardware for GPU numbers: NVIDIA via WebGPU, DP4A int8 dot path, exact-gated
5
  against the verified units at init.
6
 
@@ -16,6 +16,7 @@ against the verified units at init.
16
  | `test_optimizer.js` | DaisyAdam beats SGD through the units, deterministic replicas | PASS — 1.59 vs 1.95, replica diff 0.000e+0 |
17
  | `test_transformer.js` | the transformer LM trains through the units end to end | PASS — loss 4.75 → 1.26 (baseline 4.56), replica diff 0.000e+0 |
18
  | `test_unit_backward.js` | int8 STE gradients do not damage convergence | PASS — units/float loss ratio 1.007 |
 
19
 
20
  ## The IEEE-754 oracle (`test_ieee.js`)
21
 
 
1
  # Test results — updated 2026-07-17
2
 
3
+ All suites run with `npm test` (chains all twelve). Every suite exits 0.
4
  Hardware for GPU numbers: NVIDIA via WebGPU, DP4A int8 dot path, exact-gated
5
  against the verified units at init.
6
 
 
16
  | `test_optimizer.js` | DaisyAdam beats SGD through the units, deterministic replicas | PASS — 1.59 vs 1.95, replica diff 0.000e+0 |
17
  | `test_transformer.js` | the transformer LM trains through the units end to end | PASS — loss 4.75 → 1.26 (baseline 4.56), replica diff 0.000e+0 |
18
  | `test_unit_backward.js` | int8 STE gradients do not damage convergence | PASS — units/float loss ratio 1.007 |
19
+ | `test_nonfinite.js` | every quantizer refuses NaN/Inf instead of zeroing the layer | PASS — 13/13; the old quantizers fail 12 (added 2026-09-27) |
20
 
21
  ## The IEEE-754 oracle (`test_ieee.js`)
22
 
web/package.json CHANGED
@@ -6,7 +6,7 @@
6
  "author": "Dean Byrne (Quazim0t0) / DaisyChainAI",
7
  "scripts": {
8
  "start": "node server.js",
9
- "test": "node test_core.js && node test_verified.js && node test_ieee.js && node test_gates.js && node test_metamorphic.js && node test_corpus.js && node test_selfcorpus.js && node test_b2b.js && node test_optimizer.js && node test_transformer.js && node test_unit_backward.js"
10
  },
11
  "dependencies": {
12
  "hyparquet": "^1.26.2",
 
6
  "author": "Dean Byrne (Quazim0t0) / DaisyChainAI",
7
  "scripts": {
8
  "start": "node server.js",
9
+ "test": "node test_core.js && node test_verified.js && node test_ieee.js && node test_gates.js && node test_metamorphic.js && node test_corpus.js && node test_selfcorpus.js && node test_b2b.js && node test_optimizer.js && node test_transformer.js && node test_unit_backward.js && node test_nonfinite.js"
10
  },
11
  "dependencies": {
12
  "hyparquet": "^1.26.2",
web/public/requant_lut.bin CHANGED
@@ -1,3 +1,3 @@
1
  version https://git-lfs.github.com/spec/v1
2
- oid sha256:173444ecfa293433329a333289983a665c481d913e9fd1c2778b55380ca4dd31
3
  size 65536
 
1
  version https://git-lfs.github.com/spec/v1
2
+ oid sha256:d567c49ab3e3d7863a8b1d1af4e178d5c8eba059835348b947095be4969a93e2
3
  size 65536
web/public/transformer.js CHANGED
@@ -1,582 +1,583 @@
1
- // A miniature transformer language model that trains THROUGH the verified INT8
2
- // units: every matrix product in the forward pass — QKV projections, attention
3
- // scores, attention·values, output projection, the MLP, and the unembedding —
4
- // runs through the verified multiply LUT (an emulated INT8 tensor core).
5
- // Backward is a straight-through estimator in float (the integer path has no
6
- // gradient), exactly like the Python VerifiedLinear.
7
- //
8
- // Task: next-character prediction on a deterministic, self-generated corpus
9
- // (every peer builds the same text from the same seed — nothing to download).
10
- (function (root) {
11
- "use strict";
12
-
13
- let TC, V; // TrainCore / Verified — resolved per environment at the end
14
-
15
- function mulberry32(a) { return function () { a |= 0; a = a + 0x6D2B79F5 | 0; let t = Math.imul(a ^ a >>> 15, 1 | a); t = t + Math.imul(t ^ t >>> 7, 61 | t) ^ t; return ((t ^ t >>> 14) >>> 0) / 4294967296; }; }
16
- function randn(n, rng) { const o = new Float32Array(n); for (let i = 0; i < n; i += 2) { let u = 0, v = 0; while (u === 0) u = rng(); while (v === 0) v = rng(); const m = Math.sqrt(-2 * Math.log(u)); o[i] = m * Math.cos(2 * Math.PI * v); if (i + 1 < n) o[i + 1] = m * Math.sin(2 * Math.PI * v); } return o; }
17
-
18
- // ---- corpus: deterministic cottagecore prose, identical on every peer -----
19
- const W_ADJ = ["mossy", "golden", "amber", "quiet", "little", "misty", "sunny", "wild", "cozy", "dusty", "merry", "brave"];
20
- const W_NOUN = ["fox", "hare", "owl", "badger", "toad", "sparrow", "otter", "deer", "mushroom", "acorn", "willow", "robin", "river", "meadow", "garden", "lantern"];
21
- const W_VERB = ["naps", "sings", "wanders", "hides", "dreams", "waits", "dances", "listens", "rests", "grows"];
22
- const W_PREP = ["by", "under", "near", "beside", "beyond", "inside"];
23
- function buildCorpus() {
24
- const rng = mulberry32(20260712);
25
- const pick = (a) => a[Math.floor(rng() * a.length)];
26
- let s = "";
27
- while (s.length < 60000)
28
- s += `the ${pick(W_ADJ)} ${pick(W_NOUN)} ${pick(W_VERB)} ${pick(W_PREP)} the ${pick(W_ADJ)} ${pick(W_NOUN)}. `;
29
- return s;
30
- }
31
- // ---- tokenizer -------------------------------------------------------------
32
- // Spikewhale tokenizer (tokenizer.json): byte-level greedy longest-match
33
- // ("length-max"), ~16.5k tokens. Until it loads (or if the file is missing)
34
- // a 96-char byte-level vocab keeps the app working — but ALL devices in a
35
- // group must use the same tokenizer (the config broadcast enforces it).
36
- const FALLBACK_CHARS = [...Array(95)].map((_, i) => String.fromCharCode(32 + i)).concat(["\n"]);
37
- let tok = {
38
- name: "char-96 (fallback)",
39
- vocab: Object.fromEntries(FALLBACK_CHARS.map((c, i) => [c, i])),
40
- ids: FALLBACK_CHARS, maxLen: 1, size: FALLBACK_CHARS.length,
41
- unk: 0, specials: new Set(),
42
- };
43
- tok.unk = tok.vocab[" "];
44
- function vocabSize() { return tok.size; }
45
- function tokenizerName() { return tok.name; }
46
- function loadTokenizerData(d) { // plain {vocab, vocab_size, max_token_len}
47
- const ids = new Array(d.vocab_size);
48
- for (const [t, i] of Object.entries(d.vocab)) ids[i] = t;
49
- tok = { name: `Spikewhale length-max (${d.vocab_size} tokens)`,
50
- vocab: d.vocab, ids, maxLen: d.max_token_len || 24, size: d.vocab_size,
51
- unk: d.vocab["<unk>"] ?? 1,
52
- specials: new Set(["<pad>", "<unk>", "<bos>", "<eos>", ...(d.special_tokens || [])]) };
53
- IDS = encode(CORPUS); // re-tokenize whatever corpus is loaded
54
- return tok.name;
55
- }
56
- async function loadTokenizer(url) {
57
- const r = await fetch(url || "tokenizer.json");
58
- if (!r.ok) throw new Error(`tokenizer.json HTTP ${r.status}`);
59
- return loadTokenizerData(await r.json());
60
- }
61
- function toLatin1(s) { const b = new TextEncoder().encode(s); let o = ""; for (const x of b) o += String.fromCharCode(x); return o; }
62
- function encode(text) { // greedy longest match over bytes
63
- const s = toLatin1(text), out = [];
64
- let i = 0;
65
- while (i < s.length) {
66
- let m = null;
67
- for (let L = Math.min(tok.maxLen, s.length - i); L > 0; L--) {
68
- const sub = s.substr(i, L);
69
- if (sub in tok.vocab) { m = sub; break; }
70
- }
71
- if (m === null) { out.push(tok.unk); i++; continue; }
72
- out.push(tok.vocab[m]); i += m.length;
73
- }
74
- return Int32Array.from(out);
75
- }
76
- function decode(idArr) {
77
- let s = "";
78
- for (const id of idArr) {
79
- const t = tok.ids[id];
80
- if (t === undefined || tok.specials.has(t)) continue;
81
- s += t;
82
- }
83
- const bytes = Uint8Array.from([...s].map(c => c.charCodeAt(0)));
84
- return new TextDecoder().decode(bytes);
85
- }
86
- let CORPUS = buildCorpus();
87
- let IDS = encode(CORPUS);
88
- let DATASET = "built-in corpus";
89
-
90
- // Training text: FineWeb-Edu (10BT sample), HARDCODED as the only dataset.
91
- // The serving Space reads random slices of the parquet shards straight off
92
- // the HF CDN with range requests (see server.js /data) — no dependency on
93
- // the datasets-server rows API, which 503s routinely. Each device pulls its
94
- // own random slice (that's data parallelism — batches were always
95
- // per-device anyway). Offline or on failure the built-in corpus stays.
96
- const DEFAULT_DS = "HuggingFaceFW/fineweb-edu";
97
- async function streamDataset() { // dataset choice removed on purpose
98
- const r = await fetch("data");
99
- if (!r.ok) throw new Error(`/data HTTP ${r.status}`);
100
- const text = (await r.text()).replace(/[^\x20-\x7e\n]/g, " ");
101
- if (text.length < 10000) throw new Error("too little text returned");
102
- CORPUS = text.slice(0, 500000);
103
- IDS = encode(CORPUS);
104
- DATASET = `${DEFAULT_DS} · 10BT sample (parquet via this Space)`;
105
- return { name: DATASET, chars: CORPUS.length };
106
- }
107
- const streamFineWebEdu = () => streamDataset();
108
- function datasetName() { return DATASET; }
109
-
110
- // ---- verified matmul: block-scaled INT8 through the units ------------------
111
- // CUTLASS ex. 67/81 blockwise scaling: per-row activation scales × per-column
112
- // weight scales, one exact LUT/DP4A GEMM, dequant (+ optional fused ReLU) in
113
- // the kernel epilogue (ex. 12). Replaces the per-tensor 3-pass: same outlier
114
- // robustness at one third of the unit ops.
115
- async function vmm(Xf, Wf, m, k, n, ctx, relu) {
116
- // ctx.audit re-checks random cells of this LIVE GEMM against the units
117
- return V.vgemmBlock(Xf, Wf, { m, k, n, batch: 1, relu: !!relu }, ctx.L, ctx.bgemm, ctx.audit);
118
- }
119
- // CUTLASS ex. 45 (dual GEMM): sibling GEMMs that share the same LEFT operand
120
- // run as ONE batched dispatch, and the shared operand is quantized ONCE
121
- // instead of once per sibling. Used for the q/k/v projections — same X
122
- // (ln1.y), three weights, identical shapes. Bit-identical to three separate
123
- // vmm calls: quantizeRows is deterministic (same input -> same int8+scales),
124
- // the tiled copies index exactly like separate batch elements, and block
125
- // scales are per-row/per-column PER BATCH ELEMENT, so concatenation changes
126
- // no scale and no product. The batched kernel is the same exact-gated bgemm
127
- // that training already runs, and the live-shape audit still samples it.
128
- async function vmmShared3(Xf, Wa, Wb, Wc, m, k, n, ctx) {
129
- const x = V.quantizeRows(Xf, m, k);
130
- const xq = new Int8Array(3 * m * k), xs = new Float32Array(3 * m);
131
- for (let i = 0; i < 3; i++) { xq.set(x.q, i * m * k); xs.set(x.s, i * m); }
132
- const wq = new Int8Array(3 * k * n), ws = new Float32Array(3 * n);
133
- [Wa, Wb, Wc].forEach((W, i) => { const w = V.quantizeCols(W, k, n); wq.set(w.q, i * k * n); ws.set(w.s, i * n); });
134
- const d = { m, k, n, batch: 3 };
135
- let out;
136
- if (ctx.bgemm) {
137
- out = await ctx.bgemm(xq, wq, xs, ws, d);
138
- if (ctx.audit && ctx.audit.due()) {
139
- const bad = V.auditTile(xq, wq, xs, ws, d, out, ctx.L, ctx.audit.cells);
140
- if (bad) ctx.audit.fail(bad);
141
- }
142
- } else {
143
- out = V.bgemmJS(xq, wq, xs, ws, d, ctx.L);
144
- }
145
- const MN = m * n;
146
- return [out.subarray(0, MN), out.subarray(MN, 2 * MN), out.subarray(2 * MN, 3 * MN)];
147
- }
148
-
149
- // ---- layernorm (no affine) -------------------------------------------------
150
- function lnFwd(x, rows, C) {
151
- const y = new Float32Array(rows * C), sig = new Float32Array(rows);
152
- for (let r = 0; r < rows; r++) {
153
- let mu = 0; for (let j = 0; j < C; j++) mu += x[r * C + j]; mu /= C;
154
- let v = 0; for (let j = 0; j < C; j++) { const d = x[r * C + j] - mu; v += d * d; }
155
- const s = Math.sqrt(v / C + 1e-5); sig[r] = s;
156
- for (let j = 0; j < C; j++) y[r * C + j] = (x[r * C + j] - mu) / s;
157
- }
158
- return { y, sig };
159
- }
160
- function lnBwd(dy, y, sig, rows, C) {
161
- const dx = new Float32Array(rows * C);
162
- for (let r = 0; r < rows; r++) {
163
- let mdy = 0, mdyy = 0;
164
- for (let j = 0; j < C; j++) { mdy += dy[r * C + j]; mdyy += dy[r * C + j] * y[r * C + j]; }
165
- mdy /= C; mdyy /= C;
166
- for (let j = 0; j < C; j++) dx[r * C + j] = (dy[r * C + j] - mdy - y[r * C + j] * mdyy) / sig[r];
167
- }
168
- return dx;
169
- }
170
-
171
- // ---- model -----------------------------------------------------------------
172
- // cfg: { c: width, t: seq len, b: batch/device, layers, heads, steps, lr }
173
- // engine: the Compute backend object ({bgemm} for the fused WebGPU path) or a
174
- // legacy matmulInt8 function (Node tests, inference kit) -> CPU LUT mirror
175
- function init(cfg, L, engine, audit) {
176
- const c = cfg.c, layers = cfg.layers || 2, heads = cfg.heads || 2, hidden = 2 * c;
177
- let seed = 100;
178
- const mk = (nEl, scale) => { const w = randn(nEl, mulberry32(seed++)); for (let i = 0; i < nEl; i++) w[i] *= scale; return w; };
179
- const params = [], names = [];
180
- const add = (name, w) => { params.push(w); names.push(name); return w; };
181
- const m = {
182
- cfg: { ...cfg, layers, heads, hidden, vocab: vocabSize() },
183
- ctx: { L, bgemm: (engine && engine.bgemm) || null,
184
- att: (engine && engine.att) || null, fgemm: (engine && engine.fgemm) || null,
185
- fgemm2: (engine && engine.fgemm2) || null,
186
- mlp: (engine && engine.mlp) || null,
187
- audit: audit || null, unitBackward: !!cfg.unitBackward },
188
- emb: add("emb", mk(vocabSize() * c, 0.08)),
189
- pos: add("pos", mk(cfg.t * c, 0.02)),
190
- blocks: [], params, names,
191
- };
192
- for (let l = 0; l < layers; l++)
193
- m.blocks.push({
194
- Wq: add(`b${l}.Wq`, mk(c * c, 0.08)), Wk: add(`b${l}.Wk`, mk(c * c, 0.08)),
195
- Wv: add(`b${l}.Wv`, mk(c * c, 0.08)), Wo: add(`b${l}.Wo`, mk(c * c, 0.08)),
196
- W1: add(`b${l}.W1`, mk(c * hidden, 0.08)), W2: add(`b${l}.W2`, mk(hidden * c, 0.08)),
197
- });
198
- // weight-tied unembedding: logits use embᵀ (no separate Wu). Halves the
199
- // vocab-sized parameters — and with a 16k vocab that's ~half of ALL
200
- // parameters, so gradients over the wire shrink ~2× too.
201
- m.nParams = params.reduce((a, p) => a + p.length, 0);
202
- return m;
203
- }
204
-
205
- function sampleBatch(cfg) {
206
- const { b, t } = cfg;
207
- const X = new Int32Array(b * t), Y = new Int32Array(b * t);
208
- for (let i = 0; i < b; i++) {
209
- const off = Math.floor(Math.random() * (IDS.length - t - 1));
210
- for (let j = 0; j < t; j++) { X[i * t + j] = IDS[off + j]; Y[i * t + j] = IDS[off + j + 1]; }
211
- }
212
- return { X, Y };
213
- }
214
-
215
- // ---- head layout helpers ---------------------------------------------------
216
- // q/k/v live as BT×C with head h owning columns [h*hd, (h+1)*hd). The backward
217
- // wants every head as its own GEMM problem, so gather once into BH×T×hd and
218
- // scatter back at the end — one pass each, instead of slicing per head inside
219
- // the loop and paying a GPU dispatch per tiny matmul.
220
- function gatherHeads(x, B, T, C, heads, hd) { // BT×C -> BH×T×hd
221
- const out = new Float32Array(B * heads * T * hd);
222
- for (let bi = 0; bi < B; bi++)
223
- for (let h = 0; h < heads; h++) {
224
- const bz = bi * heads + h;
225
- for (let ti = 0; ti < T; ti++)
226
- for (let j = 0; j < hd; j++) out[(bz * T + ti) * hd + j] = x[(bi * T + ti) * C + h * hd + j];
227
- }
228
- return out;
229
- }
230
- function scatterHeadsAcc(dst, src, B, T, C, heads, hd) { // BH×T×hd -> += BT×C
231
- for (let bi = 0; bi < B; bi++)
232
- for (let h = 0; h < heads; h++) {
233
- const bz = bi * heads + h;
234
- for (let ti = 0; ti < T; ti++)
235
- for (let j = 0; j < hd; j++) dst[(bi * T + ti) * C + h * hd + j] += src[(bz * T + ti) * hd + j];
236
- }
237
- }
238
- function batchedTranspose(x, batch, rows, cols) { // per-batch rows×cols -> cols×rows
239
- const out = new Float32Array(batch * rows * cols);
240
- for (let b = 0; b < batch; b++) {
241
- const o = b * rows * cols;
242
- for (let r = 0; r < rows; r++)
243
- for (let c = 0; c < cols; c++) out[o + c * rows + r] = x[o + r * cols + c];
244
- }
245
- return out;
246
- }
247
-
248
- // ---- forward THROUGH the verified units (caches kept for STE backward) -----
249
- async function forward(m, X, Y) {
250
- const { c: C, t: T, b: B, layers, heads, hidden, vocab } = m.cfg;
251
- const BT = B * T, hd = C / heads, ctx = m.ctx;
252
- const cache = { X, Y, blocks: [] };
253
- let x = new Float32Array(BT * C);
254
- for (let i = 0; i < BT; i++) {
255
- const id = X[i], tpos = i % T;
256
- for (let j = 0; j < C; j++) x[i * C + j] = m.emb[id * C + j] + m.pos[tpos * C + j];
257
- }
258
- for (let l = 0; l < layers; l++) {
259
- const bl = m.blocks[l], cb = { xin: x };
260
- const l1 = lnFwd(x, BT, C); cb.ln1 = l1;
261
- // q/k/v share the same left operand — one batched dispatch, one quantize
262
- // of ln1.y instead of three (CUTLASS ex. 45; see vmmShared3)
263
- const [q, k, v] = await vmmShared3(l1.y, bl.Wq, bl.Wk, bl.Wv, BT, C, C, ctx);
264
- cb.q = q; cb.k = k; cb.v = v;
265
- const scale = 1 / Math.sqrt(hd);
266
- // gather-FUSED attention (CUTLASS ex. 36/52): the kernels read q/k/v in
267
- // their natural BT×C layout with head-strided indexing and scatter ctx
268
- // straight back — no JS gather copies, no kᵀ transpose. All B×H heads in
269
- // one dispatch per stage, every product through the verified units.
270
- const BH = B * heads, dAtt = { B, T, heads, hd };
271
- // per-(token,head) row quantization: the (BT·heads)×hd view IS the buffer
272
- const qq = V.quantizeRows(q, BT * heads, hd), kq = V.quantizeRows(k, BT * heads, hd);
273
- const sAll = ctx.att ? await ctx.att.scores(qq.q, kq.q, qq.s, kq.s, dAtt)
274
- : V.attScoresJS(qq.q, kq.q, qq.s, kq.s, dAtt, ctx.L);
275
- // live-shape audit: the init gate only ever saw four test shapes
276
- if (ctx.att && ctx.audit && ctx.audit.due()) {
277
- const bad = V.auditAttScores(qq.q, kq.q, qq.s, kq.s, dAtt, sAll, ctx.L, ctx.audit.cells);
278
- if (bad) ctx.audit.fail(bad);
279
- }
280
- const aAll = new Float32Array(BH * T * T); // causal softmax
281
- for (let bz = 0; bz < BH; bz++) {
282
- const so = bz * T * T;
283
- for (let ti = 0; ti < T; ti++) {
284
- let mx = -1e30;
285
- for (let tj = 0; tj <= ti; tj++) mx = Math.max(mx, sAll[so + ti * T + tj] * scale);
286
- let z = 0;
287
- for (let tj = 0; tj <= ti; tj++) { const e = Math.exp(sAll[so + ti * T + tj] * scale - mx); aAll[so + ti * T + tj] = e; z += e; }
288
- for (let tj = 0; tj <= ti; tj++) aAll[so + ti * T + tj] /= z;
289
- }
290
- }
291
- const aq = V.quantizeRows(aAll, BH * T, T);
292
- const vq = V.quantizeHeadCols(v, B, T, heads, hd);
293
- const ctxOut = ctx.att ? await ctx.att.ctx(aq.q, vq.q, aq.s, vq.s, dAtt)
294
- : V.attCtxJS(aq.q, vq.q, aq.s, vq.s, dAtt, ctx.L);
295
- if (ctx.att && ctx.audit && ctx.audit.due()) {
296
- const bad = V.auditAttCtx(aq.q, vq.q, aq.s, vq.s, dAtt, ctxOut, ctx.L, ctx.audit.cells);
297
- if (bad) ctx.audit.fail(bad);
298
- }
299
- cb.aAll = aAll; // backward slices heads from q/k/v/aAll
300
- cb.ctxOut = ctxOut;
301
- const attnOut = await vmm(ctxOut, bl.Wo, BT, C, C, ctx);
302
- const x2 = new Float32Array(BT * C);
303
- for (let i = 0; i < x2.length; i++) x2[i] = x[i] + attnOut[i];
304
- cb.x2 = x2;
305
- const l2 = lnFwd(x2, BT, C); cb.ln2 = l2;
306
- // CUTLASS ex. 13 + 23: both MLP GEMMs run back-to-back on the GPU. The
307
- // intermediate h1 is quantized ON-DEVICE (exact-gated respec — see
308
- // vmlpBlock in verified_core.js) and only its per-row absmax (~1KB)
309
- // visits JS between the GEMMs; h1 itself comes back solely because the
310
- // STE backward needs it. CPU devices run the bit-identical mirror chain.
311
- const { h1, out: mlpOut } = await V.vmlpBlock(l2.y, bl.W1, bl.W2,
312
- { m: BT, k: C, h: hidden, n: C }, ctx.L, ctx.mlp, ctx.audit);
313
- const mask = new Uint8Array(h1.length);
314
- for (let i = 0; i < h1.length; i++) if (h1[i] > 0) mask[i] = 1;
315
- cb.h1 = h1; cb.mask = mask;
316
- x = new Float32Array(BT * C);
317
- for (let i = 0; i < x.length; i++) x[i] = x2[i] + mlpOut[i];
318
- cache.blocks.push(cb);
319
- }
320
- const lf = lnFwd(x, BT, C); cache.lnf = lf; cache.xf = x;
321
- const logits = await vmm(lf.y, TC.transpose(m.emb, vocab, C), BT, C, vocab, ctx); // tied: embᵀ
322
- // cross-entropy + dlogits
323
- let loss = 0;
324
- const dlogits = new Float32Array(BT * vocab);
325
- for (let i = 0; i < BT; i++) {
326
- let mx = -1e30;
327
- for (let j = 0; j < vocab; j++) mx = Math.max(mx, logits[i * vocab + j]);
328
- let z = 0;
329
- for (let j = 0; j < vocab; j++) z += Math.exp(logits[i * vocab + j] - mx);
330
- const lz = Math.log(z) + mx;
331
- loss += lz - logits[i * vocab + Y[i]];
332
- for (let j = 0; j < vocab; j++)
333
- dlogits[i * vocab + j] = (Math.exp(logits[i * vocab + j] - lz) - (j === Y[i] ? 1 : 0)) / BT;
334
- }
335
- loss /= BT;
336
- cache.dlogits = dlogits;
337
- return { loss, cache, logits };
338
- }
339
-
340
- // ---- STE backward (float), mirrors forward exactly --------------------------
341
- // The two vocab-sized matmuls run on the split-K f32 GPU kernel when
342
- // available (CUTLASS ex. 06) — same float math, off the JS thread.
343
- async function backward(m, cache) {
344
- const { c: C, t: T, b: B, layers, heads, hidden, vocab } = m.cfg;
345
- const BT = B * T, hd = C / heads, tr = TC.transpose;
346
- const g = m.params.map(p => new Float32Array(p.length));
347
- const gi = Object.fromEntries(m.names.map((n, i) => [n, i]));
348
- // Every matmul here goes through `bmm`. With ctx.unitBackward the STE
349
- // gradient is computed BY the verified units (block-scaled int8, exact int32
350
- // accumulate) instead of in float. STE is a claim about the math — pretend
351
- // the quantizer was the identity — not about the datatype that evaluates it,
352
- // so the two are orthogonal and this stays a correct STE.
353
- const units = !!m.ctx.unitBackward;
354
- const bmm = units
355
- ? (A, Bm, mm_, k, n) => vmm(A, Bm, mm_, k, n, m.ctx)
356
- : async (A, Bm, mm_, k, n) => TC.matmul(A, Bm, mm_, k, n);
357
- // batched: all `batch` problems in ONE dispatch (CUTLASS ex. 05/24). The
358
- // per-head backward is 4 GEMMs x B x heads of tiny matrices; issued one at a
359
- // time the GPU spends all its time on dispatch overhead rather than math.
360
- const bmmB = units
361
- ? (A, Bm, rows, k, n, batch) =>
362
- V.vgemmBlock(A, Bm, { m: rows, k, n, batch }, m.ctx.L, m.ctx.bgemm, m.ctx.audit)
363
- : async (A, Bm, rows, k, n, batch) => {
364
- const out = new Float32Array(batch * rows * n);
365
- for (let bz = 0; bz < batch; bz++)
366
- out.set(TC.matmul(A.subarray(bz * rows * k, (bz + 1) * rows * k),
367
- Bm.subarray(bz * k * n, (bz + 1) * k * n), rows, k, n), bz * rows * n);
368
- return out;
369
- };
370
- // tied unembed: logits = lnf @ embᵀ, so the unembedding gradient flows
371
- // straight into emb — dlogitsᵀ @ lnf is V×C, emb's own shape
372
- let dlnfIn;
373
- if (m.ctx.fgemm2 && !units) {
374
- // Both GEMMs consume dlogits (BT x vocab, ~17 MB at the 16512 vocab).
375
- // fgemm2 uploads it ONCE and runs both on one submit — profiling had
376
- // this pair at 55% of the step, over half of it re-uploading the same
377
- // operand. Bit-identical to the two separate calls (gated at init).
378
- [g[gi.emb], dlnfIn] = await m.ctx.fgemm2(
379
- cache.dlogits,
380
- cache.lnf.y, { m: vocab, k: BT, n: C, transA: true },
381
- m.emb, { m: BT, k: vocab, n: C }); // split-K shape
382
- } else if (m.ctx.fgemm && !units) {
383
- [g[gi.emb], dlnfIn] = await Promise.all([
384
- m.ctx.fgemm(cache.dlogits, cache.lnf.y, { m: vocab, k: BT, n: C, transA: true }),
385
- m.ctx.fgemm(cache.dlogits, m.emb, { m: BT, k: vocab, n: C }), // split-K shape
386
- ]);
387
- } else if (units && m.ctx.bgemm) {
388
- // units + GPU: the g.emb operand is dlogitsᵀ (vocab×BT, ~4M elements), and
389
- // tr() + quantizeRows() is three full passes over it in JS. Quantizing the
390
- // COLUMNS of dlogits directly into transposed int8 is one pass and
391
- // bit-identical: same |max| scan, same rounds, in the same order — only
392
- // the write pattern changes. The GEMM itself still goes through ctx.bgemm
393
- // (exact-gated), and the live-shape audit still samples it.
394
- const quantizeColsAsRows = (X, rows, cols) => { // == quantizeRows(tr(X), cols, rows)
395
- const q = new Int8Array(cols * rows), s = new Float32Array(cols);
396
- for (let c = 0; c < cols; c++) {
397
- let mx = 0;
398
- for (let r = 0; r < rows; r++) { const a = Math.abs(X[r * cols + c]); if (a > mx) mx = a; }
399
- const sc = Math.max(mx / 127, 1e-8); s[c] = sc;
400
- for (let r = 0; r < rows; r++) {
401
- const v = Math.round(X[r * cols + c] / sc);
402
- q[c * rows + r] = v < -128 ? -128 : v > 127 ? 127 : v;
403
- }
404
- }
405
- return { q, s };
406
- };
407
- const dlq = quantizeColsAsRows(cache.dlogits, BT, vocab); // dlogitsᵀ quantized, one pass
408
- const wq2 = V.quantizeCols(cache.lnf.y, BT, C);
409
- const dEmb = { m: vocab, k: BT, n: C, batch: 1 };
410
- const [gEmb, dIn] = await Promise.all([
411
- m.ctx.bgemm(dlq.q, wq2.q, dlq.s, wq2.s, dEmb),
412
- bmm(cache.dlogits, m.emb, BT, vocab, C),
413
- ]);
414
- if (m.ctx.audit && m.ctx.audit.due()) {
415
- const bad = V.auditTile(dlq.q, wq2.q, dlq.s, wq2.s, dEmb, gEmb, m.ctx.L, m.ctx.audit.cells);
416
- if (bad) m.ctx.audit.fail(bad);
417
- }
418
- g[gi.emb] = gEmb; dlnfIn = dIn;
419
- } else {
420
- // independent GEMMs — overlap them (these are the two vocab-sized calls,
421
- // the largest in the whole backward; each is its own round trip)
422
- [g[gi.emb], dlnfIn] = await Promise.all([
423
- bmm(tr(cache.dlogits, BT, vocab), cache.lnf.y, vocab, BT, C),
424
- bmm(cache.dlogits, m.emb, BT, vocab, C),
425
- ]);
426
- }
427
- let dx = lnBwd(dlnfIn, cache.lnf.y, cache.lnf.sig, BT, C);
428
- const scale = 1 / Math.sqrt(hd);
429
- // concat helper for fusing sibling GEMMs into one batched dispatch
430
- const cat = (...arrs) => {
431
- const out = new Float32Array(arrs.reduce((a, x) => a + x.length, 0));
432
- let o = 0; for (const x of arrs) { out.set(x, o); o += x.length; }
433
- return out;
434
- };
435
- for (let l = layers - 1; l >= 0; l--) {
436
- const bl = m.blocks[l], cb = cache.blocks[l];
437
- // mlp: x3 = x2 + relu(ln2 @ W1) @ W2
438
- // gW2 and dh1 are independent — overlap their dispatches. On GPU each bmm
439
- // is a full upload/submit/readback round trip, so sequential awaits leave
440
- // the GPU idle between every pair; this is pure latency, not arithmetic,
441
- // and each GEMM's int32 accumulation is exact so overlap changes no bit.
442
- const dmlpOut = dx; // residual passthrough handled below
443
- const [gW2, dh1] = await Promise.all([
444
- bmm(tr(cb.h1, BT, hidden), dmlpOut, hidden, BT, C),
445
- bmm(dmlpOut, tr(bl.W2, hidden, C), BT, C, hidden),
446
- ]);
447
- g[gi[`b${l}.W2`]] = gW2;
448
- for (let i = 0; i < dh1.length; i++) if (!cb.mask[i]) dh1[i] = 0;
449
- const [gW1, dln2raw] = await Promise.all([
450
- bmm(tr(cb.ln2.y, BT, C), dh1, C, BT, hidden),
451
- bmm(dh1, tr(bl.W1, C, hidden), BT, hidden, C),
452
- ]);
453
- g[gi[`b${l}.W1`]] = gW1;
454
- const dln2in = lnBwd(dln2raw, cb.ln2.y, cb.ln2.sig, BT, C);
455
- const dx2 = new Float32Array(BT * C);
456
- for (let i = 0; i < dx2.length; i++) dx2[i] = dx[i] + dln2in[i];
457
- // attention: x2 = xin + (ctxOut @ Wo)
458
- const [gWo, dctx] = await Promise.all([
459
- bmm(tr(cb.ctxOut, BT, C), dx2, C, BT, C),
460
- bmm(dx2, tr(bl.Wo, C, C), BT, C, C),
461
- ]);
462
- g[gi[`b${l}.Wo`]] = gWo;
463
- // gather every head once, then run each stage as ONE batched GEMM over all
464
- // B*heads problems: 4 dispatches per layer instead of 4 per head.
465
- const BH = B * heads;
466
- const qb = gatherHeads(cb.q, B, T, C, heads, hd);
467
- const kb = gatherHeads(cb.k, B, T, C, heads, hd);
468
- const vb = gatherHeads(cb.v, B, T, C, heads, hd);
469
- const dchb = gatherHeads(dctx, B, T, C, heads, hd);
470
- const aT = batchedTranspose(cb.aAll, BH, T, T); // BH×T×T
471
- const vT = batchedTranspose(vb, BH, T, hd); // BH×hd×T
472
- const [dvAll, daAll] = await Promise.all([
473
- bmmB(aT, dchb, T, T, hd, BH), // aᵀ @ dctx
474
- bmmB(dchb, vT, T, hd, T, BH), // dctx @ vᵀ
475
- ]);
476
- // softmax backward is elementwise + a causal row reduction: stays in float
477
- // (no matrix math here, so nothing for the units to do)
478
- const dsAll = new Float32Array(BH * T * T);
479
- for (let bz = 0; bz < BH; bz++) {
480
- const o = bz * T * T;
481
- for (let ti = 0; ti < T; ti++) {
482
- let dot = 0;
483
- for (let tj = 0; tj <= ti; tj++) dot += daAll[o + ti * T + tj] * cb.aAll[o + ti * T + tj];
484
- for (let tj = 0; tj <= ti; tj++)
485
- dsAll[o + ti * T + tj] = cb.aAll[o + ti * T + tj] * (daAll[o + ti * T + tj] - dot) * scale;
486
- }
487
- }
488
- const dsT = batchedTranspose(dsAll, BH, T, T);
489
- const [dqAll, dkAll] = await Promise.all([
490
- bmmB(dsAll, kb, T, T, hd, BH), // ds @ k
491
- bmmB(dsT, qb, T, T, hd, BH), // dsᵀ @ q
492
- ]);
493
- const dq = new Float32Array(BT * C), dk = new Float32Array(BT * C), dv = new Float32Array(BT * C);
494
- scatterHeadsAcc(dq, dqAll, B, T, C, heads, hd);
495
- scatterHeadsAcc(dk, dkAll, B, T, C, heads, hd);
496
- scatterHeadsAcc(dv, dvAll, B, T, C, heads, hd);
497
- // The QKV weight grads share the same left operand (ln1ᵀ), and the three
498
- // dln1in terms share one shape — each trio fuses into ONE batched GEMM
499
- // (batch=3) instead of three dispatches. Bit-identical to separate calls:
500
- // block scales are per-row of X and per-column of W PER BATCH ELEMENT, so
501
- // concatenation changes no scale and no product.
502
- const ln1T = tr(cb.ln1.y, BT, C);
503
- const [gQKV, dIn3] = await Promise.all([
504
- bmmB(cat(ln1T, ln1T, ln1T), cat(dq, dk, dv), C, BT, C, 3),
505
- bmmB(cat(dq, dk, dv), cat(tr(bl.Wq, C, C), tr(bl.Wk, C, C), tr(bl.Wv, C, C)), BT, C, C, 3),
506
- ]);
507
- const CC = C * C, BTC = BT * C;
508
- g[gi[`b${l}.Wq`]] = gQKV.slice(0, CC);
509
- g[gi[`b${l}.Wk`]] = gQKV.slice(CC, 2 * CC);
510
- g[gi[`b${l}.Wv`]] = gQKV.slice(2 * CC, 3 * CC);
511
- // sum the three dln1in terms in q,k,v order with an f32 round after EACH
512
- // add — the old code accumulated into a Float32Array element three times,
513
- // which rounds per step; a bare q+k+v here would run in f64 and round
514
- // once, a last-ulp difference that forks replicas. (Exactly the epilogue
515
- // mirror lesson: match the rounding schedule, not just the values.)
516
- const dln1in = new Float32Array(BTC);
517
- for (let i = 0; i < BTC; i++)
518
- dln1in[i] = Math.fround(Math.fround(dIn3[i] + dIn3[BTC + i]) + dIn3[2 * BTC + i]);
519
- const dxin = lnBwd(dln1in, cb.ln1.y, cb.ln1.sig, BT, C);
520
- dx = new Float32Array(BT * C);
521
- for (let i = 0; i < dx.length; i++) dx[i] = dx2[i] + dxin[i];
522
- }
523
- // embedding + positional
524
- const ge = g[gi.emb], gp = g[gi.pos];
525
- for (let i = 0; i < BT; i++) {
526
- const id = cache.X[i], tpos = i % T;
527
- for (let j = 0; j < C; j++) { ge[id * C + j] += dx[i * C + j]; gp[tpos * C + j] += dx[i * C + j]; }
528
- }
529
- // flatten
530
- const flat = new Float32Array(m.nParams);
531
- let off = 0;
532
- for (const t of g) { flat.set(t, off); off += t.length; }
533
- return flat;
534
- }
535
-
536
- async function trainStep(m) {
537
- const { X, Y } = sampleBatch(m.cfg);
538
- const { loss, cache } = await forward(m, X, Y);
539
- const grad = await backward(m, cache);
540
- return { loss, grad };
541
- }
542
-
543
- function applyUpdate(m, upd) { // W -= upd (lr folded in by the optimizer)
544
- let off = 0;
545
- for (const p of m.params) { for (let i = 0; i < p.length; i++) p[i] -= upd[off + i]; off += p.length; }
546
- }
547
- function getFlatParams(m) {
548
- const flat = new Float32Array(m.nParams);
549
- let off = 0;
550
- for (const p of m.params) { flat.set(p, off); off += p.length; }
551
- return flat;
552
- }
553
- function setFlatParams(m, flat) {
554
- let off = 0;
555
- for (const p of m.params) { p.set(flat.subarray(off, off + p.length)); off += p.length; }
556
- }
557
-
558
- // greedy sampling — watch the model actually speak
559
- async function generate(m, prompt, nChars) {
560
- const { t: T } = m.cfg;
561
- let ids = [...encode(prompt)];
562
- for (let step = 0; step < nChars; step++) {
563
- const win = ids.slice(-T);
564
- const X = new Int32Array(T), Y = new Int32Array(T);
565
- for (let i = 0; i < win.length; i++) X[T - win.length + i] = win[i];
566
- const save = m.cfg.b; m.cfg.b = 1;
567
- const { logits } = await forward(m, X, Y);
568
- m.cfg.b = save;
569
- const row = (T - 1) * m.cfg.vocab;
570
- let best = 0, bv = -1e30;
571
- for (let j = 0; j < m.cfg.vocab; j++) if (logits[row + j] > bv) { bv = logits[row + j]; best = j; }
572
- ids.push(best);
573
- }
574
- return decode(ids);
575
- }
576
-
577
- const api = { init, trainStep, applyUpdate, getFlatParams, setFlatParams, generate,
578
- streamFineWebEdu, streamDataset, datasetName, loadTokenizer, loadTokenizerData,
579
- vocabSize, tokenizerName, encode, decode };
580
- if (typeof module !== "undefined" && module.exports) { TC = require("./traincore.js"); V = require("./verified_core.js"); module.exports = api; }
581
- else { TC = root.TrainCore; V = root.Verified; root.Transformer = api; }
582
- })(typeof self !== "undefined" ? self : this);
 
 
1
+ // A miniature transformer language model that trains THROUGH the verified INT8
2
+ // units: every matrix product in the forward pass — QKV projections, attention
3
+ // scores, attention·values, output projection, the MLP, and the unembedding —
4
+ // runs through the verified multiply LUT (an emulated INT8 tensor core).
5
+ // Backward is a straight-through estimator in float (the integer path has no
6
+ // gradient), exactly like the Python VerifiedLinear.
7
+ //
8
+ // Task: next-character prediction on a deterministic, self-generated corpus
9
+ // (every peer builds the same text from the same seed — nothing to download).
10
+ (function (root) {
11
+ "use strict";
12
+
13
+ let TC, V; // TrainCore / Verified — resolved per environment at the end
14
+
15
+ function mulberry32(a) { return function () { a |= 0; a = a + 0x6D2B79F5 | 0; let t = Math.imul(a ^ a >>> 15, 1 | a); t = t + Math.imul(t ^ t >>> 7, 61 | t) ^ t; return ((t ^ t >>> 14) >>> 0) / 4294967296; }; }
16
+ function randn(n, rng) { const o = new Float32Array(n); for (let i = 0; i < n; i += 2) { let u = 0, v = 0; while (u === 0) u = rng(); while (v === 0) v = rng(); const m = Math.sqrt(-2 * Math.log(u)); o[i] = m * Math.cos(2 * Math.PI * v); if (i + 1 < n) o[i + 1] = m * Math.sin(2 * Math.PI * v); } return o; }
17
+
18
+ // ---- corpus: deterministic cottagecore prose, identical on every peer -----
19
+ const W_ADJ = ["mossy", "golden", "amber", "quiet", "little", "misty", "sunny", "wild", "cozy", "dusty", "merry", "brave"];
20
+ const W_NOUN = ["fox", "hare", "owl", "badger", "toad", "sparrow", "otter", "deer", "mushroom", "acorn", "willow", "robin", "river", "meadow", "garden", "lantern"];
21
+ const W_VERB = ["naps", "sings", "wanders", "hides", "dreams", "waits", "dances", "listens", "rests", "grows"];
22
+ const W_PREP = ["by", "under", "near", "beside", "beyond", "inside"];
23
+ function buildCorpus() {
24
+ const rng = mulberry32(20260712);
25
+ const pick = (a) => a[Math.floor(rng() * a.length)];
26
+ let s = "";
27
+ while (s.length < 60000)
28
+ s += `the ${pick(W_ADJ)} ${pick(W_NOUN)} ${pick(W_VERB)} ${pick(W_PREP)} the ${pick(W_ADJ)} ${pick(W_NOUN)}. `;
29
+ return s;
30
+ }
31
+ // ---- tokenizer -------------------------------------------------------------
32
+ // Spikewhale tokenizer (tokenizer.json): byte-level greedy longest-match
33
+ // ("length-max"), ~16.5k tokens. Until it loads (or if the file is missing)
34
+ // a 96-char byte-level vocab keeps the app working — but ALL devices in a
35
+ // group must use the same tokenizer (the config broadcast enforces it).
36
+ const FALLBACK_CHARS = [...Array(95)].map((_, i) => String.fromCharCode(32 + i)).concat(["\n"]);
37
+ let tok = {
38
+ name: "char-96 (fallback)",
39
+ vocab: Object.fromEntries(FALLBACK_CHARS.map((c, i) => [c, i])),
40
+ ids: FALLBACK_CHARS, maxLen: 1, size: FALLBACK_CHARS.length,
41
+ unk: 0, specials: new Set(),
42
+ };
43
+ tok.unk = tok.vocab[" "];
44
+ function vocabSize() { return tok.size; }
45
+ function tokenizerName() { return tok.name; }
46
+ function loadTokenizerData(d) { // plain {vocab, vocab_size, max_token_len}
47
+ const ids = new Array(d.vocab_size);
48
+ for (const [t, i] of Object.entries(d.vocab)) ids[i] = t;
49
+ tok = { name: `Spikewhale length-max (${d.vocab_size} tokens)`,
50
+ vocab: d.vocab, ids, maxLen: d.max_token_len || 24, size: d.vocab_size,
51
+ unk: d.vocab["<unk>"] ?? 1,
52
+ specials: new Set(["<pad>", "<unk>", "<bos>", "<eos>", ...(d.special_tokens || [])]) };
53
+ IDS = encode(CORPUS); // re-tokenize whatever corpus is loaded
54
+ return tok.name;
55
+ }
56
+ async function loadTokenizer(url) {
57
+ const r = await fetch(url || "tokenizer.json");
58
+ if (!r.ok) throw new Error(`tokenizer.json HTTP ${r.status}`);
59
+ return loadTokenizerData(await r.json());
60
+ }
61
+ function toLatin1(s) { const b = new TextEncoder().encode(s); let o = ""; for (const x of b) o += String.fromCharCode(x); return o; }
62
+ function encode(text) { // greedy longest match over bytes
63
+ const s = toLatin1(text), out = [];
64
+ let i = 0;
65
+ while (i < s.length) {
66
+ let m = null;
67
+ for (let L = Math.min(tok.maxLen, s.length - i); L > 0; L--) {
68
+ const sub = s.substr(i, L);
69
+ if (sub in tok.vocab) { m = sub; break; }
70
+ }
71
+ if (m === null) { out.push(tok.unk); i++; continue; }
72
+ out.push(tok.vocab[m]); i += m.length;
73
+ }
74
+ return Int32Array.from(out);
75
+ }
76
+ function decode(idArr) {
77
+ let s = "";
78
+ for (const id of idArr) {
79
+ const t = tok.ids[id];
80
+ if (t === undefined || tok.specials.has(t)) continue;
81
+ s += t;
82
+ }
83
+ const bytes = Uint8Array.from([...s].map(c => c.charCodeAt(0)));
84
+ return new TextDecoder().decode(bytes);
85
+ }
86
+ let CORPUS = buildCorpus();
87
+ let IDS = encode(CORPUS);
88
+ let DATASET = "built-in corpus";
89
+
90
+ // Training text: FineWeb-Edu (10BT sample), HARDCODED as the only dataset.
91
+ // The serving Space reads random slices of the parquet shards straight off
92
+ // the HF CDN with range requests (see server.js /data) — no dependency on
93
+ // the datasets-server rows API, which 503s routinely. Each device pulls its
94
+ // own random slice (that's data parallelism — batches were always
95
+ // per-device anyway). Offline or on failure the built-in corpus stays.
96
+ const DEFAULT_DS = "HuggingFaceFW/fineweb-edu";
97
+ async function streamDataset() { // dataset choice removed on purpose
98
+ const r = await fetch("data");
99
+ if (!r.ok) throw new Error(`/data HTTP ${r.status}`);
100
+ const text = (await r.text()).replace(/[^\x20-\x7e\n]/g, " ");
101
+ if (text.length < 10000) throw new Error("too little text returned");
102
+ CORPUS = text.slice(0, 500000);
103
+ IDS = encode(CORPUS);
104
+ DATASET = `${DEFAULT_DS} · 10BT sample (parquet via this Space)`;
105
+ return { name: DATASET, chars: CORPUS.length };
106
+ }
107
+ const streamFineWebEdu = () => streamDataset();
108
+ function datasetName() { return DATASET; }
109
+
110
+ // ---- verified matmul: block-scaled INT8 through the units ------------------
111
+ // CUTLASS ex. 67/81 blockwise scaling: per-row activation scales × per-column
112
+ // weight scales, one exact LUT/DP4A GEMM, dequant (+ optional fused ReLU) in
113
+ // the kernel epilogue (ex. 12). Replaces the per-tensor 3-pass: same outlier
114
+ // robustness at one third of the unit ops.
115
+ async function vmm(Xf, Wf, m, k, n, ctx, relu) {
116
+ // ctx.audit re-checks random cells of this LIVE GEMM against the units
117
+ return V.vgemmBlock(Xf, Wf, { m, k, n, batch: 1, relu: !!relu }, ctx.L, ctx.bgemm, ctx.audit);
118
+ }
119
+ // CUTLASS ex. 45 (dual GEMM): sibling GEMMs that share the same LEFT operand
120
+ // run as ONE batched dispatch, and the shared operand is quantized ONCE
121
+ // instead of once per sibling. Used for the q/k/v projections — same X
122
+ // (ln1.y), three weights, identical shapes. Bit-identical to three separate
123
+ // vmm calls: quantizeRows is deterministic (same input -> same int8+scales),
124
+ // the tiled copies index exactly like separate batch elements, and block
125
+ // scales are per-row/per-column PER BATCH ELEMENT, so concatenation changes
126
+ // no scale and no product. The batched kernel is the same exact-gated bgemm
127
+ // that training already runs, and the live-shape audit still samples it.
128
+ async function vmmShared3(Xf, Wa, Wb, Wc, m, k, n, ctx) {
129
+ const x = V.quantizeRows(Xf, m, k);
130
+ const xq = new Int8Array(3 * m * k), xs = new Float32Array(3 * m);
131
+ for (let i = 0; i < 3; i++) { xq.set(x.q, i * m * k); xs.set(x.s, i * m); }
132
+ const wq = new Int8Array(3 * k * n), ws = new Float32Array(3 * n);
133
+ [Wa, Wb, Wc].forEach((W, i) => { const w = V.quantizeCols(W, k, n); wq.set(w.q, i * k * n); ws.set(w.s, i * n); });
134
+ const d = { m, k, n, batch: 3 };
135
+ let out;
136
+ if (ctx.bgemm) {
137
+ out = await ctx.bgemm(xq, wq, xs, ws, d);
138
+ if (ctx.audit && ctx.audit.due()) {
139
+ const bad = V.auditTile(xq, wq, xs, ws, d, out, ctx.L, ctx.audit.cells);
140
+ if (bad) ctx.audit.fail(bad);
141
+ }
142
+ } else {
143
+ out = V.bgemmJS(xq, wq, xs, ws, d, ctx.L);
144
+ }
145
+ const MN = m * n;
146
+ return [out.subarray(0, MN), out.subarray(MN, 2 * MN), out.subarray(2 * MN, 3 * MN)];
147
+ }
148
+
149
+ // ---- layernorm (no affine) -------------------------------------------------
150
+ function lnFwd(x, rows, C) {
151
+ const y = new Float32Array(rows * C), sig = new Float32Array(rows);
152
+ for (let r = 0; r < rows; r++) {
153
+ let mu = 0; for (let j = 0; j < C; j++) mu += x[r * C + j]; mu /= C;
154
+ let v = 0; for (let j = 0; j < C; j++) { const d = x[r * C + j] - mu; v += d * d; }
155
+ const s = Math.sqrt(v / C + 1e-5); sig[r] = s;
156
+ for (let j = 0; j < C; j++) y[r * C + j] = (x[r * C + j] - mu) / s;
157
+ }
158
+ return { y, sig };
159
+ }
160
+ function lnBwd(dy, y, sig, rows, C) {
161
+ const dx = new Float32Array(rows * C);
162
+ for (let r = 0; r < rows; r++) {
163
+ let mdy = 0, mdyy = 0;
164
+ for (let j = 0; j < C; j++) { mdy += dy[r * C + j]; mdyy += dy[r * C + j] * y[r * C + j]; }
165
+ mdy /= C; mdyy /= C;
166
+ for (let j = 0; j < C; j++) dx[r * C + j] = (dy[r * C + j] - mdy - y[r * C + j] * mdyy) / sig[r];
167
+ }
168
+ return dx;
169
+ }
170
+
171
+ // ---- model -----------------------------------------------------------------
172
+ // cfg: { c: width, t: seq len, b: batch/device, layers, heads, steps, lr }
173
+ // engine: the Compute backend object ({bgemm} for the fused WebGPU path) or a
174
+ // legacy matmulInt8 function (Node tests, inference kit) -> CPU LUT mirror
175
+ function init(cfg, L, engine, audit) {
176
+ const c = cfg.c, layers = cfg.layers || 2, heads = cfg.heads || 2, hidden = 2 * c;
177
+ let seed = 100;
178
+ const mk = (nEl, scale) => { const w = randn(nEl, mulberry32(seed++)); for (let i = 0; i < nEl; i++) w[i] *= scale; return w; };
179
+ const params = [], names = [];
180
+ const add = (name, w) => { params.push(w); names.push(name); return w; };
181
+ const m = {
182
+ cfg: { ...cfg, layers, heads, hidden, vocab: vocabSize() },
183
+ ctx: { L, bgemm: (engine && engine.bgemm) || null,
184
+ att: (engine && engine.att) || null, fgemm: (engine && engine.fgemm) || null,
185
+ fgemm2: (engine && engine.fgemm2) || null,
186
+ mlp: (engine && engine.mlp) || null,
187
+ audit: audit || null, unitBackward: !!cfg.unitBackward },
188
+ emb: add("emb", mk(vocabSize() * c, 0.08)),
189
+ pos: add("pos", mk(cfg.t * c, 0.02)),
190
+ blocks: [], params, names,
191
+ };
192
+ for (let l = 0; l < layers; l++)
193
+ m.blocks.push({
194
+ Wq: add(`b${l}.Wq`, mk(c * c, 0.08)), Wk: add(`b${l}.Wk`, mk(c * c, 0.08)),
195
+ Wv: add(`b${l}.Wv`, mk(c * c, 0.08)), Wo: add(`b${l}.Wo`, mk(c * c, 0.08)),
196
+ W1: add(`b${l}.W1`, mk(c * hidden, 0.08)), W2: add(`b${l}.W2`, mk(hidden * c, 0.08)),
197
+ });
198
+ // weight-tied unembedding: logits use embᵀ (no separate Wu). Halves the
199
+ // vocab-sized parameters — and with a 16k vocab that's ~half of ALL
200
+ // parameters, so gradients over the wire shrink ~2× too.
201
+ m.nParams = params.reduce((a, p) => a + p.length, 0);
202
+ return m;
203
+ }
204
+
205
+ function sampleBatch(cfg) {
206
+ const { b, t } = cfg;
207
+ const X = new Int32Array(b * t), Y = new Int32Array(b * t);
208
+ for (let i = 0; i < b; i++) {
209
+ const off = Math.floor(Math.random() * (IDS.length - t - 1));
210
+ for (let j = 0; j < t; j++) { X[i * t + j] = IDS[off + j]; Y[i * t + j] = IDS[off + j + 1]; }
211
+ }
212
+ return { X, Y };
213
+ }
214
+
215
+ // ---- head layout helpers ---------------------------------------------------
216
+ // q/k/v live as BT×C with head h owning columns [h*hd, (h+1)*hd). The backward
217
+ // wants every head as its own GEMM problem, so gather once into BH×T×hd and
218
+ // scatter back at the end — one pass each, instead of slicing per head inside
219
+ // the loop and paying a GPU dispatch per tiny matmul.
220
+ function gatherHeads(x, B, T, C, heads, hd) { // BT×C -> BH×T×hd
221
+ const out = new Float32Array(B * heads * T * hd);
222
+ for (let bi = 0; bi < B; bi++)
223
+ for (let h = 0; h < heads; h++) {
224
+ const bz = bi * heads + h;
225
+ for (let ti = 0; ti < T; ti++)
226
+ for (let j = 0; j < hd; j++) out[(bz * T + ti) * hd + j] = x[(bi * T + ti) * C + h * hd + j];
227
+ }
228
+ return out;
229
+ }
230
+ function scatterHeadsAcc(dst, src, B, T, C, heads, hd) { // BH×T×hd -> += BT×C
231
+ for (let bi = 0; bi < B; bi++)
232
+ for (let h = 0; h < heads; h++) {
233
+ const bz = bi * heads + h;
234
+ for (let ti = 0; ti < T; ti++)
235
+ for (let j = 0; j < hd; j++) dst[(bi * T + ti) * C + h * hd + j] += src[(bz * T + ti) * hd + j];
236
+ }
237
+ }
238
+ function batchedTranspose(x, batch, rows, cols) { // per-batch rows×cols -> cols×rows
239
+ const out = new Float32Array(batch * rows * cols);
240
+ for (let b = 0; b < batch; b++) {
241
+ const o = b * rows * cols;
242
+ for (let r = 0; r < rows; r++)
243
+ for (let c = 0; c < cols; c++) out[o + c * rows + r] = x[o + r * cols + c];
244
+ }
245
+ return out;
246
+ }
247
+
248
+ // ---- forward THROUGH the verified units (caches kept for STE backward) -----
249
+ async function forward(m, X, Y) {
250
+ const { c: C, t: T, b: B, layers, heads, hidden, vocab } = m.cfg;
251
+ const BT = B * T, hd = C / heads, ctx = m.ctx;
252
+ const cache = { X, Y, blocks: [] };
253
+ let x = new Float32Array(BT * C);
254
+ for (let i = 0; i < BT; i++) {
255
+ const id = X[i], tpos = i % T;
256
+ for (let j = 0; j < C; j++) x[i * C + j] = m.emb[id * C + j] + m.pos[tpos * C + j];
257
+ }
258
+ for (let l = 0; l < layers; l++) {
259
+ const bl = m.blocks[l], cb = { xin: x };
260
+ const l1 = lnFwd(x, BT, C); cb.ln1 = l1;
261
+ // q/k/v share the same left operand — one batched dispatch, one quantize
262
+ // of ln1.y instead of three (CUTLASS ex. 45; see vmmShared3)
263
+ const [q, k, v] = await vmmShared3(l1.y, bl.Wq, bl.Wk, bl.Wv, BT, C, C, ctx);
264
+ cb.q = q; cb.k = k; cb.v = v;
265
+ const scale = 1 / Math.sqrt(hd);
266
+ // gather-FUSED attention (CUTLASS ex. 36/52): the kernels read q/k/v in
267
+ // their natural BT×C layout with head-strided indexing and scatter ctx
268
+ // straight back — no JS gather copies, no kᵀ transpose. All B×H heads in
269
+ // one dispatch per stage, every product through the verified units.
270
+ const BH = B * heads, dAtt = { B, T, heads, hd };
271
+ // per-(token,head) row quantization: the (BT·heads)×hd view IS the buffer
272
+ const qq = V.quantizeRows(q, BT * heads, hd), kq = V.quantizeRows(k, BT * heads, hd);
273
+ const sAll = ctx.att ? await ctx.att.scores(qq.q, kq.q, qq.s, kq.s, dAtt)
274
+ : V.attScoresJS(qq.q, kq.q, qq.s, kq.s, dAtt, ctx.L);
275
+ // live-shape audit: the init gate only ever saw four test shapes
276
+ if (ctx.att && ctx.audit && ctx.audit.due()) {
277
+ const bad = V.auditAttScores(qq.q, kq.q, qq.s, kq.s, dAtt, sAll, ctx.L, ctx.audit.cells);
278
+ if (bad) ctx.audit.fail(bad);
279
+ }
280
+ const aAll = new Float32Array(BH * T * T); // causal softmax
281
+ for (let bz = 0; bz < BH; bz++) {
282
+ const so = bz * T * T;
283
+ for (let ti = 0; ti < T; ti++) {
284
+ let mx = -1e30;
285
+ for (let tj = 0; tj <= ti; tj++) mx = Math.max(mx, sAll[so + ti * T + tj] * scale);
286
+ let z = 0;
287
+ for (let tj = 0; tj <= ti; tj++) { const e = Math.exp(sAll[so + ti * T + tj] * scale - mx); aAll[so + ti * T + tj] = e; z += e; }
288
+ for (let tj = 0; tj <= ti; tj++) aAll[so + ti * T + tj] /= z;
289
+ }
290
+ }
291
+ const aq = V.quantizeRows(aAll, BH * T, T);
292
+ const vq = V.quantizeHeadCols(v, B, T, heads, hd);
293
+ const ctxOut = ctx.att ? await ctx.att.ctx(aq.q, vq.q, aq.s, vq.s, dAtt)
294
+ : V.attCtxJS(aq.q, vq.q, aq.s, vq.s, dAtt, ctx.L);
295
+ if (ctx.att && ctx.audit && ctx.audit.due()) {
296
+ const bad = V.auditAttCtx(aq.q, vq.q, aq.s, vq.s, dAtt, ctxOut, ctx.L, ctx.audit.cells);
297
+ if (bad) ctx.audit.fail(bad);
298
+ }
299
+ cb.aAll = aAll; // backward slices heads from q/k/v/aAll
300
+ cb.ctxOut = ctxOut;
301
+ const attnOut = await vmm(ctxOut, bl.Wo, BT, C, C, ctx);
302
+ const x2 = new Float32Array(BT * C);
303
+ for (let i = 0; i < x2.length; i++) x2[i] = x[i] + attnOut[i];
304
+ cb.x2 = x2;
305
+ const l2 = lnFwd(x2, BT, C); cb.ln2 = l2;
306
+ // CUTLASS ex. 13 + 23: both MLP GEMMs run back-to-back on the GPU. The
307
+ // intermediate h1 is quantized ON-DEVICE (exact-gated respec — see
308
+ // vmlpBlock in verified_core.js) and only its per-row absmax (~1KB)
309
+ // visits JS between the GEMMs; h1 itself comes back solely because the
310
+ // STE backward needs it. CPU devices run the bit-identical mirror chain.
311
+ const { h1, out: mlpOut } = await V.vmlpBlock(l2.y, bl.W1, bl.W2,
312
+ { m: BT, k: C, h: hidden, n: C }, ctx.L, ctx.mlp, ctx.audit);
313
+ const mask = new Uint8Array(h1.length);
314
+ for (let i = 0; i < h1.length; i++) if (h1[i] > 0) mask[i] = 1;
315
+ cb.h1 = h1; cb.mask = mask;
316
+ x = new Float32Array(BT * C);
317
+ for (let i = 0; i < x.length; i++) x[i] = x2[i] + mlpOut[i];
318
+ cache.blocks.push(cb);
319
+ }
320
+ const lf = lnFwd(x, BT, C); cache.lnf = lf; cache.xf = x;
321
+ const logits = await vmm(lf.y, TC.transpose(m.emb, vocab, C), BT, C, vocab, ctx); // tied: embᵀ
322
+ // cross-entropy + dlogits
323
+ let loss = 0;
324
+ const dlogits = new Float32Array(BT * vocab);
325
+ for (let i = 0; i < BT; i++) {
326
+ let mx = -1e30;
327
+ for (let j = 0; j < vocab; j++) mx = Math.max(mx, logits[i * vocab + j]);
328
+ let z = 0;
329
+ for (let j = 0; j < vocab; j++) z += Math.exp(logits[i * vocab + j] - mx);
330
+ const lz = Math.log(z) + mx;
331
+ loss += lz - logits[i * vocab + Y[i]];
332
+ for (let j = 0; j < vocab; j++)
333
+ dlogits[i * vocab + j] = (Math.exp(logits[i * vocab + j] - lz) - (j === Y[i] ? 1 : 0)) / BT;
334
+ }
335
+ loss /= BT;
336
+ cache.dlogits = dlogits;
337
+ return { loss, cache, logits };
338
+ }
339
+
340
+ // ---- STE backward (float), mirrors forward exactly --------------------------
341
+ // The two vocab-sized matmuls run on the split-K f32 GPU kernel when
342
+ // available (CUTLASS ex. 06) — same float math, off the JS thread.
343
+ async function backward(m, cache) {
344
+ const { c: C, t: T, b: B, layers, heads, hidden, vocab } = m.cfg;
345
+ const BT = B * T, hd = C / heads, tr = TC.transpose;
346
+ const g = m.params.map(p => new Float32Array(p.length));
347
+ const gi = Object.fromEntries(m.names.map((n, i) => [n, i]));
348
+ // Every matmul here goes through `bmm`. With ctx.unitBackward the STE
349
+ // gradient is computed BY the verified units (block-scaled int8, exact int32
350
+ // accumulate) instead of in float. STE is a claim about the math — pretend
351
+ // the quantizer was the identity — not about the datatype that evaluates it,
352
+ // so the two are orthogonal and this stays a correct STE.
353
+ const units = !!m.ctx.unitBackward;
354
+ const bmm = units
355
+ ? (A, Bm, mm_, k, n) => vmm(A, Bm, mm_, k, n, m.ctx)
356
+ : async (A, Bm, mm_, k, n) => TC.matmul(A, Bm, mm_, k, n);
357
+ // batched: all `batch` problems in ONE dispatch (CUTLASS ex. 05/24). The
358
+ // per-head backward is 4 GEMMs x B x heads of tiny matrices; issued one at a
359
+ // time the GPU spends all its time on dispatch overhead rather than math.
360
+ const bmmB = units
361
+ ? (A, Bm, rows, k, n, batch) =>
362
+ V.vgemmBlock(A, Bm, { m: rows, k, n, batch }, m.ctx.L, m.ctx.bgemm, m.ctx.audit)
363
+ : async (A, Bm, rows, k, n, batch) => {
364
+ const out = new Float32Array(batch * rows * n);
365
+ for (let bz = 0; bz < batch; bz++)
366
+ out.set(TC.matmul(A.subarray(bz * rows * k, (bz + 1) * rows * k),
367
+ Bm.subarray(bz * k * n, (bz + 1) * k * n), rows, k, n), bz * rows * n);
368
+ return out;
369
+ };
370
+ // tied unembed: logits = lnf @ embᵀ, so the unembedding gradient flows
371
+ // straight into emb — dlogitsᵀ @ lnf is V×C, emb's own shape
372
+ let dlnfIn;
373
+ if (m.ctx.fgemm2 && !units) {
374
+ // Both GEMMs consume dlogits (BT x vocab, ~17 MB at the 16512 vocab).
375
+ // fgemm2 uploads it ONCE and runs both on one submit — profiling had
376
+ // this pair at 55% of the step, over half of it re-uploading the same
377
+ // operand. Bit-identical to the two separate calls (gated at init).
378
+ [g[gi.emb], dlnfIn] = await m.ctx.fgemm2(
379
+ cache.dlogits,
380
+ cache.lnf.y, { m: vocab, k: BT, n: C, transA: true },
381
+ m.emb, { m: BT, k: vocab, n: C }); // split-K shape
382
+ } else if (m.ctx.fgemm && !units) {
383
+ [g[gi.emb], dlnfIn] = await Promise.all([
384
+ m.ctx.fgemm(cache.dlogits, cache.lnf.y, { m: vocab, k: BT, n: C, transA: true }),
385
+ m.ctx.fgemm(cache.dlogits, m.emb, { m: BT, k: vocab, n: C }), // split-K shape
386
+ ]);
387
+ } else if (units && m.ctx.bgemm) {
388
+ // units + GPU: the g.emb operand is dlogitsᵀ (vocab×BT, ~4M elements), and
389
+ // tr() + quantizeRows() is three full passes over it in JS. Quantizing the
390
+ // COLUMNS of dlogits directly into transposed int8 is one pass and
391
+ // bit-identical: same |max| scan, same rounds, in the same order — only
392
+ // the write pattern changes. The GEMM itself still goes through ctx.bgemm
393
+ // (exact-gated), and the live-shape audit still samples it.
394
+ const quantizeColsAsRows = (X, rows, cols) => { // == quantizeRows(tr(X), cols, rows)
395
+ V.assertFinite(X, "quantizeColsAsRows");
396
+ const q = new Int8Array(cols * rows), s = new Float32Array(cols);
397
+ for (let c = 0; c < cols; c++) {
398
+ let mx = 0;
399
+ for (let r = 0; r < rows; r++) { const a = Math.abs(X[r * cols + c]); if (a > mx) mx = a; }
400
+ const sc = Math.max(mx / 127, 1e-8); s[c] = sc;
401
+ for (let r = 0; r < rows; r++) {
402
+ const v = Math.round(X[r * cols + c] / sc);
403
+ q[c * rows + r] = v < -128 ? -128 : v > 127 ? 127 : v;
404
+ }
405
+ }
406
+ return { q, s };
407
+ };
408
+ const dlq = quantizeColsAsRows(cache.dlogits, BT, vocab); // dlogitsᵀ quantized, one pass
409
+ const wq2 = V.quantizeCols(cache.lnf.y, BT, C);
410
+ const dEmb = { m: vocab, k: BT, n: C, batch: 1 };
411
+ const [gEmb, dIn] = await Promise.all([
412
+ m.ctx.bgemm(dlq.q, wq2.q, dlq.s, wq2.s, dEmb),
413
+ bmm(cache.dlogits, m.emb, BT, vocab, C),
414
+ ]);
415
+ if (m.ctx.audit && m.ctx.audit.due()) {
416
+ const bad = V.auditTile(dlq.q, wq2.q, dlq.s, wq2.s, dEmb, gEmb, m.ctx.L, m.ctx.audit.cells);
417
+ if (bad) m.ctx.audit.fail(bad);
418
+ }
419
+ g[gi.emb] = gEmb; dlnfIn = dIn;
420
+ } else {
421
+ // independent GEMMs — overlap them (these are the two vocab-sized calls,
422
+ // the largest in the whole backward; each is its own round trip)
423
+ [g[gi.emb], dlnfIn] = await Promise.all([
424
+ bmm(tr(cache.dlogits, BT, vocab), cache.lnf.y, vocab, BT, C),
425
+ bmm(cache.dlogits, m.emb, BT, vocab, C),
426
+ ]);
427
+ }
428
+ let dx = lnBwd(dlnfIn, cache.lnf.y, cache.lnf.sig, BT, C);
429
+ const scale = 1 / Math.sqrt(hd);
430
+ // concat helper for fusing sibling GEMMs into one batched dispatch
431
+ const cat = (...arrs) => {
432
+ const out = new Float32Array(arrs.reduce((a, x) => a + x.length, 0));
433
+ let o = 0; for (const x of arrs) { out.set(x, o); o += x.length; }
434
+ return out;
435
+ };
436
+ for (let l = layers - 1; l >= 0; l--) {
437
+ const bl = m.blocks[l], cb = cache.blocks[l];
438
+ // mlp: x3 = x2 + relu(ln2 @ W1) @ W2
439
+ // gW2 and dh1 are independent — overlap their dispatches. On GPU each bmm
440
+ // is a full upload/submit/readback round trip, so sequential awaits leave
441
+ // the GPU idle between every pair; this is pure latency, not arithmetic,
442
+ // and each GEMM's int32 accumulation is exact so overlap changes no bit.
443
+ const dmlpOut = dx; // residual passthrough handled below
444
+ const [gW2, dh1] = await Promise.all([
445
+ bmm(tr(cb.h1, BT, hidden), dmlpOut, hidden, BT, C),
446
+ bmm(dmlpOut, tr(bl.W2, hidden, C), BT, C, hidden),
447
+ ]);
448
+ g[gi[`b${l}.W2`]] = gW2;
449
+ for (let i = 0; i < dh1.length; i++) if (!cb.mask[i]) dh1[i] = 0;
450
+ const [gW1, dln2raw] = await Promise.all([
451
+ bmm(tr(cb.ln2.y, BT, C), dh1, C, BT, hidden),
452
+ bmm(dh1, tr(bl.W1, C, hidden), BT, hidden, C),
453
+ ]);
454
+ g[gi[`b${l}.W1`]] = gW1;
455
+ const dln2in = lnBwd(dln2raw, cb.ln2.y, cb.ln2.sig, BT, C);
456
+ const dx2 = new Float32Array(BT * C);
457
+ for (let i = 0; i < dx2.length; i++) dx2[i] = dx[i] + dln2in[i];
458
+ // attention: x2 = xin + (ctxOut @ Wo)
459
+ const [gWo, dctx] = await Promise.all([
460
+ bmm(tr(cb.ctxOut, BT, C), dx2, C, BT, C),
461
+ bmm(dx2, tr(bl.Wo, C, C), BT, C, C),
462
+ ]);
463
+ g[gi[`b${l}.Wo`]] = gWo;
464
+ // gather every head once, then run each stage as ONE batched GEMM over all
465
+ // B*heads problems: 4 dispatches per layer instead of 4 per head.
466
+ const BH = B * heads;
467
+ const qb = gatherHeads(cb.q, B, T, C, heads, hd);
468
+ const kb = gatherHeads(cb.k, B, T, C, heads, hd);
469
+ const vb = gatherHeads(cb.v, B, T, C, heads, hd);
470
+ const dchb = gatherHeads(dctx, B, T, C, heads, hd);
471
+ const aT = batchedTranspose(cb.aAll, BH, T, T); // BH×T×T
472
+ const vT = batchedTranspose(vb, BH, T, hd); // BH×hd×T
473
+ const [dvAll, daAll] = await Promise.all([
474
+ bmmB(aT, dchb, T, T, hd, BH), // aᵀ @ dctx
475
+ bmmB(dchb, vT, T, hd, T, BH), // dctx @ vᵀ
476
+ ]);
477
+ // softmax backward is elementwise + a causal row reduction: stays in float
478
+ // (no matrix math here, so nothing for the units to do)
479
+ const dsAll = new Float32Array(BH * T * T);
480
+ for (let bz = 0; bz < BH; bz++) {
481
+ const o = bz * T * T;
482
+ for (let ti = 0; ti < T; ti++) {
483
+ let dot = 0;
484
+ for (let tj = 0; tj <= ti; tj++) dot += daAll[o + ti * T + tj] * cb.aAll[o + ti * T + tj];
485
+ for (let tj = 0; tj <= ti; tj++)
486
+ dsAll[o + ti * T + tj] = cb.aAll[o + ti * T + tj] * (daAll[o + ti * T + tj] - dot) * scale;
487
+ }
488
+ }
489
+ const dsT = batchedTranspose(dsAll, BH, T, T);
490
+ const [dqAll, dkAll] = await Promise.all([
491
+ bmmB(dsAll, kb, T, T, hd, BH), // ds @ k
492
+ bmmB(dsT, qb, T, T, hd, BH), // dsᵀ @ q
493
+ ]);
494
+ const dq = new Float32Array(BT * C), dk = new Float32Array(BT * C), dv = new Float32Array(BT * C);
495
+ scatterHeadsAcc(dq, dqAll, B, T, C, heads, hd);
496
+ scatterHeadsAcc(dk, dkAll, B, T, C, heads, hd);
497
+ scatterHeadsAcc(dv, dvAll, B, T, C, heads, hd);
498
+ // The QKV weight grads share the same left operand (ln1ᵀ), and the three
499
+ // dln1in terms share one shape — each trio fuses into ONE batched GEMM
500
+ // (batch=3) instead of three dispatches. Bit-identical to separate calls:
501
+ // block scales are per-row of X and per-column of W PER BATCH ELEMENT, so
502
+ // concatenation changes no scale and no product.
503
+ const ln1T = tr(cb.ln1.y, BT, C);
504
+ const [gQKV, dIn3] = await Promise.all([
505
+ bmmB(cat(ln1T, ln1T, ln1T), cat(dq, dk, dv), C, BT, C, 3),
506
+ bmmB(cat(dq, dk, dv), cat(tr(bl.Wq, C, C), tr(bl.Wk, C, C), tr(bl.Wv, C, C)), BT, C, C, 3),
507
+ ]);
508
+ const CC = C * C, BTC = BT * C;
509
+ g[gi[`b${l}.Wq`]] = gQKV.slice(0, CC);
510
+ g[gi[`b${l}.Wk`]] = gQKV.slice(CC, 2 * CC);
511
+ g[gi[`b${l}.Wv`]] = gQKV.slice(2 * CC, 3 * CC);
512
+ // sum the three dln1in terms in q,k,v order with an f32 round after EACH
513
+ // add — the old code accumulated into a Float32Array element three times,
514
+ // which rounds per step; a bare q+k+v here would run in f64 and round
515
+ // once, a last-ulp difference that forks replicas. (Exactly the epilogue
516
+ // mirror lesson: match the rounding schedule, not just the values.)
517
+ const dln1in = new Float32Array(BTC);
518
+ for (let i = 0; i < BTC; i++)
519
+ dln1in[i] = Math.fround(Math.fround(dIn3[i] + dIn3[BTC + i]) + dIn3[2 * BTC + i]);
520
+ const dxin = lnBwd(dln1in, cb.ln1.y, cb.ln1.sig, BT, C);
521
+ dx = new Float32Array(BT * C);
522
+ for (let i = 0; i < dx.length; i++) dx[i] = dx2[i] + dxin[i];
523
+ }
524
+ // embedding + positional
525
+ const ge = g[gi.emb], gp = g[gi.pos];
526
+ for (let i = 0; i < BT; i++) {
527
+ const id = cache.X[i], tpos = i % T;
528
+ for (let j = 0; j < C; j++) { ge[id * C + j] += dx[i * C + j]; gp[tpos * C + j] += dx[i * C + j]; }
529
+ }
530
+ // flatten
531
+ const flat = new Float32Array(m.nParams);
532
+ let off = 0;
533
+ for (const t of g) { flat.set(t, off); off += t.length; }
534
+ return flat;
535
+ }
536
+
537
+ async function trainStep(m) {
538
+ const { X, Y } = sampleBatch(m.cfg);
539
+ const { loss, cache } = await forward(m, X, Y);
540
+ const grad = await backward(m, cache);
541
+ return { loss, grad };
542
+ }
543
+
544
+ function applyUpdate(m, upd) { // W -= upd (lr folded in by the optimizer)
545
+ let off = 0;
546
+ for (const p of m.params) { for (let i = 0; i < p.length; i++) p[i] -= upd[off + i]; off += p.length; }
547
+ }
548
+ function getFlatParams(m) {
549
+ const flat = new Float32Array(m.nParams);
550
+ let off = 0;
551
+ for (const p of m.params) { flat.set(p, off); off += p.length; }
552
+ return flat;
553
+ }
554
+ function setFlatParams(m, flat) {
555
+ let off = 0;
556
+ for (const p of m.params) { p.set(flat.subarray(off, off + p.length)); off += p.length; }
557
+ }
558
+
559
+ // greedy sampling — watch the model actually speak
560
+ async function generate(m, prompt, nChars) {
561
+ const { t: T } = m.cfg;
562
+ let ids = [...encode(prompt)];
563
+ for (let step = 0; step < nChars; step++) {
564
+ const win = ids.slice(-T);
565
+ const X = new Int32Array(T), Y = new Int32Array(T);
566
+ for (let i = 0; i < win.length; i++) X[T - win.length + i] = win[i];
567
+ const save = m.cfg.b; m.cfg.b = 1;
568
+ const { logits } = await forward(m, X, Y);
569
+ m.cfg.b = save;
570
+ const row = (T - 1) * m.cfg.vocab;
571
+ let best = 0, bv = -1e30;
572
+ for (let j = 0; j < m.cfg.vocab; j++) if (logits[row + j] > bv) { bv = logits[row + j]; best = j; }
573
+ ids.push(best);
574
+ }
575
+ return decode(ids);
576
+ }
577
+
578
+ const api = { init, trainStep, applyUpdate, getFlatParams, setFlatParams, generate,
579
+ streamFineWebEdu, streamDataset, datasetName, loadTokenizer, loadTokenizerData,
580
+ vocabSize, tokenizerName, encode, decode };
581
+ if (typeof module !== "undefined" && module.exports) { TC = require("./traincore.js"); V = require("./verified_core.js"); module.exports = api; }
582
+ else { TC = root.TrainCore; V = root.Verified; root.Transformer = api; }
583
+ })(typeof self !== "undefined" ? self : this);
web/public/verified_core.js CHANGED
@@ -1,496 +1,519 @@
1
- // Verified INT8 compute — the emulated GPU logic, in the browser.
2
- // A layer's forward runs THROUGH the units: quantize -> LUT multiply -> requant
3
- // -> optional ReLU -> dequant. Backward is a straight-through estimator (the
4
- // integer path has no gradient), so ordinary float weights still learn.
5
- // Same units as the Python/Docker DaisyChain; here they're lookup tables.
6
- (function (root) {
7
- "use strict";
8
-
9
- let TC; // TrainCore (matmul/transpose) — resolved per environment at the end
10
-
11
- function quantize(X) {
12
- let mx = 0; for (let i = 0; i < X.length; i++) { const a = Math.abs(X[i]); if (a > mx) mx = a; }
13
- const scale = Math.max(mx / 127, 1e-8);
14
- const q = new Int8Array(X.length);
15
- for (let i = 0; i < X.length; i++) { let v = Math.round(X[i] / scale); q[i] = v < -128 ? -128 : v > 127 ? 127 : v; }
16
- return { q, scale };
17
- }
18
-
19
- // int8 matmul via the verified multiply LUT: acc(m×n) = sum_k mulLUT[Xq,Wq]
20
- function lutMatmulJS(Xq, Wq, m, k, n, L) {
21
- const C = new Int32Array(m * n), mul = L.mul;
22
- for (let i = 0; i < m; i++) {
23
- for (let p = 0; p < k; p++) {
24
- const au = (Xq[i * k + p] & 0xFF) * 256, wo = p * n, co = i * n;
25
- for (let j = 0; j < n; j++) C[co + j] += mul[au + (Wq[wo + j] & 0xFF)];
26
- }
27
- }
28
- return C;
29
- }
30
-
31
- // ---- 3xINT8 fast-accurate GEMM --------------------------------------------
32
- // The CUTLASS example-27 "3xTF32" scheme, ported to the verified units:
33
- // split each float into a coarse int8 part plus an int8-quantized residual,
34
- // run three EXACT LUT GEMMs (hi·hi, hi·lo, lo·hi), drop the negligible
35
- // lo·lo, and recombine. Same big/small decomposition NVIDIA uses to recover
36
- // near-fp32 accuracy from TF32 tensor cores — here it recovers ~14-bit
37
- // accuracy from the 8-bit units, at 3× the unit ops. Every product still
38
- // goes through the verified mul8 LUT.
39
- function quantize2(X) {
40
- const hi = quantize(X);
41
- const r = new Float32Array(X.length);
42
- for (let i = 0; i < X.length; i++) r[i] = X[i] - hi.q[i] * hi.scale;
43
- const lo = quantize(r);
44
- return { hi, lo };
45
- }
46
- function combine3(hh, hl, lh, x, w, len) {
47
- const out = new Float32Array(len);
48
- const shh = x.hi.scale * w.hi.scale, shl = x.hi.scale * w.lo.scale, slh = x.lo.scale * w.hi.scale;
49
- for (let i = 0; i < len; i++) out[i] = hh[i] * shh + hl[i] * shl + lh[i] * slh;
50
- return out;
51
- }
52
- function lutMatmul3JS(Xf, Wf, m, k, n, L) { // sync, CPU LUT path
53
- const x = quantize2(Xf), w = quantize2(Wf);
54
- return combine3(lutMatmulJS(x.hi.q, w.hi.q, m, k, n, L),
55
- lutMatmulJS(x.hi.q, w.lo.q, m, k, n, L),
56
- lutMatmulJS(x.lo.q, w.hi.q, m, k, n, L), x, w, m * n);
57
- }
58
- async function lutMatmul3(Xf, Wf, m, k, n, L, matmulInt8) { // any backend
59
- const x = quantize2(Xf), w = quantize2(Wf);
60
- const mm = matmulInt8 || lutMatmulJS;
61
- const [hh, hl, lh] = await Promise.all([
62
- mm(x.hi.q, w.hi.q, m, k, n, L),
63
- mm(x.hi.q, w.lo.q, m, k, n, L),
64
- mm(x.lo.q, w.hi.q, m, k, n, L),
65
- ]);
66
- return combine3(hh, hl, lh, x, w, m * n);
67
- }
68
-
69
- // ---- block-scaled verified GEMM (CUTLASS ex. 67/81 blockwise scaling) ------
70
- // Per-ROW scales for the activations and per-COLUMN scales for the weights:
71
- // the integer math through the mul8 LUT is completely unchanged — only the
72
- // dequant uses rs[row]·cs[col] instead of one tensor-wide product, so a single
73
- // outlier no longer crushes the quantization resolution of every other
74
- // row/column. One LUT pass at this granularity beats the per-tensor 3-pass.
75
- function quantizeRows(X, rows, cols) {
76
- const q = new Int8Array(rows * cols), s = new Float32Array(rows);
77
- for (let r = 0; r < rows; r++) {
78
- let mx = 0;
79
- for (let c = 0; c < cols; c++) { const a = Math.abs(X[r * cols + c]); if (a > mx) mx = a; }
80
- const sc = Math.max(mx / 127, 1e-8); s[r] = sc;
81
- for (let c = 0; c < cols; c++) {
82
- const v = Math.round(X[r * cols + c] / sc);
83
- q[r * cols + c] = v < -128 ? -128 : v > 127 ? 127 : v;
84
- }
85
- }
86
- return { q, s };
87
- }
88
- function quantizeCols(W, rows, cols) {
89
- const q = new Int8Array(rows * cols), s = new Float32Array(cols);
90
- for (let c = 0; c < cols; c++) {
91
- let mx = 0;
92
- for (let r = 0; r < rows; r++) { const a = Math.abs(W[r * cols + c]); if (a > mx) mx = a; }
93
- s[c] = Math.max(mx / 127, 1e-8);
94
- }
95
- for (let r = 0; r < rows; r++)
96
- for (let c = 0; c < cols; c++) {
97
- const v = Math.round(W[r * cols + c] / s[c]);
98
- q[r * cols + c] = v < -128 ? -128 : v > 127 ? 127 : v;
99
- }
100
- return { q, s };
101
- }
102
-
103
- // ---- B2B MLP chain (CUTLASS ex. 13 two-GEMM fusion + ex. 23 epilogue
104
- // reduction), respecced for cross-device exactness ---------------------------
105
- // The MLP is the one back-to-back GEMM pair with no layernorm/softmax between
106
- // (ReLU is already fused in the epilogue), so the intermediate h1 can be
107
- // quantized ON the GPU and fed straight to the second GEMM. Two rules make
108
- // that fleet-safe:
109
- // 1. The per-row |max| (ex. 23) uses only comparisons — exact on any
110
- // hardware, order-independent — and comes back to JS as ~1KB.
111
- // 2. Scale DERIVATION (two divisions) stays in JS f64, which IEEE requires
112
- // to be exactly rounded and is therefore identical on every device.
113
- // WGSL division is only 2.5 ULP — a fork waiting to happen — but WGSL
114
- // multiply/add are correctly rounded and floor/clamp are exact. So the
115
- // quantize step is respecced from round(x / scale) to
116
- // floor(f32(x * invScale) + 0.5) — floor(x+0.5) IS Math.round's tie
117
- // rule — and the fround-stepped mirror below is bit-identical to the
118
- // GPU kernel, which is exact-gated against it at init.
119
- // NOTE this changes which int8 a value on a rounding boundary lands on
120
- // (≤1 step) vs quantizeRows, so old and new builds cannot co-train — the
121
- // divergence guard stops such mixed groups by design.
122
- function rowAbsMax(X, rows, cols) {
123
- const mx = new Float32Array(rows);
124
- for (let r = 0; r < rows; r++) {
125
- let m = 0;
126
- for (let c = 0; c < cols; c++) { const a = Math.abs(X[r * cols + c]); if (a > m) m = a; }
127
- mx[r] = m;
128
- }
129
- return mx;
130
- }
131
- function scalesFromAbsMax(mx) { // f64 divisions: exactly rounded, device-identical
132
- const scale = new Float32Array(mx.length), inv = new Float32Array(mx.length);
133
- for (let i = 0; i < mx.length; i++) {
134
- scale[i] = Math.max(mx[i] / 127, 1e-8);
135
- inv[i] = 1 / scale[i]; // recip of the STORED f32 scale
136
- }
137
- return { scale, inv };
138
- }
139
- function quantizeRowsInv(X, rows, cols, inv) { // bit-exact mirror of the GPU quantize kernel
140
- const q = new Int8Array(rows * cols);
141
- for (let r = 0; r < rows; r++) {
142
- const iv = inv[r];
143
- for (let c = 0; c < cols; c++) {
144
- const n = Math.floor(f32(f32(X[r * cols + c] * iv) + 0.5));
145
- q[r * cols + c] = n < -128 ? -128 : n > 127 ? 127 : n;
146
- }
147
- }
148
- return q;
149
- }
150
- // the chained MLP: X @ W1 -> ReLU (fused) -> absmax -> quantize -> @ W2.
151
- // d = { m, k, h, n }; gpuMlp (from webgpu.js) runs both GEMMs + the
152
- // on-GPU quantize with one tiny absmax readback between; without it the CPU
153
- // mirror chain runs — SAME math, so mixed GPU/CPU fleets stay bit-identical.
154
- async function vmlpBlock(Xf, W1f, W2f, d, L, gpuMlp, audit) {
155
- const x = quantizeRows(Xf, d.m, d.k);
156
- const w1 = quantizeCols(W1f, d.k, d.h);
157
- const w2 = quantizeCols(W2f, d.h, d.n);
158
- if (gpuMlp) {
159
- const r = await gpuMlp(x.q, w1.q, w2.q, x.s, w1.s, w2.s, d);
160
- if (audit && audit.due()) {
161
- // audit BOTH live GEMMs: gemm1 against the units directly; gemm2 by
162
- // reconstructing its exact operand through the proven quantize mirror
163
- const bad1 = auditTile(x.q, w1.q, x.s, w1.s, { m: d.m, k: d.k, n: d.h, relu: true }, r.h1, L, audit.cells);
164
- if (bad1) { audit.fail("mlp gemm1: " + bad1); return r; }
165
- const sc = scalesFromAbsMax(rowAbsMax(r.h1, d.m, d.h));
166
- const hq = quantizeRowsInv(r.h1, d.m, d.h, sc.inv);
167
- const bad2 = auditTile(hq, w2.q, sc.scale, w2.s, { m: d.m, k: d.h, n: d.n }, r.out, L, audit.cells);
168
- if (bad2) audit.fail("mlp gemm2: " + bad2);
169
- }
170
- return r;
171
- }
172
- const h1 = bgemmJS(x.q, w1.q, x.s, w1.s, { m: d.m, k: d.k, n: d.h, batch: 1, relu: true }, L);
173
- const sc = scalesFromAbsMax(rowAbsMax(h1, d.m, d.h));
174
- const hq = quantizeRowsInv(h1, d.m, d.h, sc.inv);
175
- const out = bgemmJS(hq, w2.q, sc.scale, w2.s, { m: d.m, k: d.h, n: d.n, batch: 1 }, L);
176
- return { h1, out };
177
- }
178
-
179
- // ---- epilogue mirror -------------------------------------------------------
180
- // BIT-EXACT mirror of the WGSL epilogue `f32(s) * a * b`. WGSL rounds to f32
181
- // after the int->float conversion and after EACH multiply; plain JS would do
182
- // the whole chain in f64 and round once, which differs in the last ulp. That
183
- // last-ulp gap is what used to force a tolerance into the kernel gates —
184
- // mirroring the rounding exactly is what lets the gates compare with `!==`.
185
- const f32 = Math.fround;
186
- function epi(s, a, b) { return f32(f32(f32(s) * a) * b); }
187
-
188
- // f32 equality at the BIT level: `!==` says -0 === 0, but replicas are
189
- // compared by hashing raw bytes, so an audit that can't see the sign of zero
190
- // could pass a device that later forks the fleet. (Real ISAs have non-IEEE
191
- // modes that flush -0 to +0 — e.g. RDNA2 output modifiers / legacy muls.)
192
- const _fb = new Float32Array(1), _ub = new Uint32Array(_fb.buffer);
193
- function bitDiff(a, b) { _fb[0] = a; const u = _ub[0]; _fb[0] = b; return u !== _ub[0]; }
194
-
195
- // CPU mirror of the fused GPU kernel: batched int8 GEMM through the LUT with
196
- // the epilogue (block dequant + optional ReLU) applied before returning —
197
- // exactly what the WGSL kernel does on-device. d.acc=true returns the raw
198
- // int32 accumulator instead (the exact oracle the fused kernel normally hides).
199
- function bgemmJS(Xq, Wq, rs, cs, d, L) {
200
- const { m, k, n } = d, batch = d.batch || 1, relu = !!d.relu, mul = L.mul;
201
- const raw = !!d.acc;
202
- const out = raw ? new Int32Array(batch * m * n) : new Float32Array(batch * m * n);
203
- const acc = new Int32Array(n);
204
- for (let bz = 0; bz < batch; bz++) {
205
- const xo = bz * m * k, wo = bz * k * n, oo = bz * m * n, co = bz * n;
206
- for (let i = 0; i < m; i++) {
207
- acc.fill(0);
208
- const xrow = xo + i * k;
209
- for (let p = 0; p < k; p++) {
210
- const au = (Xq[xrow + p] & 0xFF) * 256, wrow = wo + p * n;
211
- for (let j = 0; j < n; j++) acc[j] += mul[au + (Wq[wrow + j] & 0xFF)];
212
- }
213
- const orow = oo + i * n;
214
- if (raw) { for (let j = 0; j < n; j++) out[orow + j] = acc[j]; continue; }
215
- const rscale = rs[bz * m + i];
216
- for (let j = 0; j < n; j++) {
217
- const v = epi(acc[j], rscale, cs[co + j]);
218
- out[orow + j] = relu && v < 0 ? 0 : v;
219
- }
220
- }
221
- }
222
- return out;
223
- }
224
-
225
- // Recompute a handful of RANDOM output cells of a live GEMM through the LUT
226
- // mirror and compare against what the kernel produced. Sampling cells instead
227
- // of whole matrices makes this cheap enough to run continuously, at the real
228
- // shapes training uses — not once at boot on toy inputs.
229
- // STRATIFIED sampling. Uniformly random cells are the wrong instrument for
230
- // the bugs that actually occur here: a bounds-guard off-by-one or a pack-tail
231
- // padding bug lives on the LAST row/column, and uniform sampling finds that
232
- // with probability ~1/n per cell — at the 16512-wide logits GEMM, never. So
233
- // the first cells are the structurally dangerous ones (corners, last row,
234
- // last column, last batch) chosen deterministically, and the remainder are
235
- // random interior cells that catch diffuse bugs. Same principle as poisoning
236
- // the buffer pool: construct the dangerous case, don't wait to land on it.
237
- function auditTile(Xq, Wq, rs, cs, d, got, L, nCells) {
238
- const { m, k, n } = d, batch = d.batch || 1, relu = !!d.relu, mul = L.mul;
239
- const N = nCells || 8;
240
- const edges = [[0, m - 1, n - 1], [0, 0, n - 1], [0, m - 1, 0], [0, 0, 0],
241
- [batch - 1, m - 1, n - 1], [batch - 1, 0, 0]];
242
- for (let t = 0; t < N; t++) {
243
- let bz, i, j;
244
- if (t < edges.length) { bz = edges[t][0]; i = edges[t][1]; j = edges[t][2]; }
245
- else { bz = (Math.random() * batch) | 0; i = (Math.random() * m) | 0; j = (Math.random() * n) | 0; }
246
- let acc = 0;
247
- const xrow = bz * m * k + i * k, wo = bz * k * n;
248
- for (let p = 0; p < k; p++) acc += mul[(Xq[xrow + p] & 0xFF) * 256 + (Wq[wo + p * n + j] & 0xFF)];
249
- let v = epi(acc, rs[bz * m + i], cs[bz * n + j]);
250
- if (relu && v < 0) v = 0;
251
- const idx = (bz * m + i) * n + j;
252
- if (bitDiff(got[idx], v))
253
- return `GEMM audit failed at [b${bz},${i},${j}] shape ${m}x${k}x${n}: kernel ${Object.is(got[idx], -0) ? "-0" : got[idx]} vs units ${Object.is(v, -0) ? "-0" : v}`;
254
- }
255
- return null;
256
- }
257
-
258
- // ---- exact mirror of the split-K f32 GEMM ----------------------------------
259
- // The f32 backward GEMM was the last kernel gated by a TOLERANCE (allclose at
260
- // 1e-3) — and this project's own gate mutation test shows allclose waving
261
- // through real bugs. The reason was real though: split-K accumulates in a
262
- // different ORDER than a naive reference, so bit-equality against the naive
263
- // one is impossible. The fix is the same as the epilogue mirror: reproduce
264
- // the kernel's order exactly, then compare with `!==`.
265
- // partials: for z in 0..S-1, sum p in [z*ks, min(k,(z+1)*ks)) in order
266
- // reduce: sum the S partials in ascending z
267
- // `fma` selects the rounding schedule for `s + a*b`: WGSL PERMITS a compiler
268
- // to contract that into a fused multiply-add (one rounding) instead of two.
269
- // Which one the device does is a fact about the device, so the gate tries
270
- // both and reports which matches rather than assuming.
271
- function fgemmMirror(A, Bm, d, fma) {
272
- const { m, k, n } = d, transA = !!d.transA;
273
- const S = k > 4096 ? Math.min(16, Math.ceil(k / 2048)) : 1;
274
- const ks = Math.ceil(k / S);
275
- const out = new Float32Array(m * n);
276
- for (let row = 0; row < m; row++)
277
- for (let col = 0; col < n; col++) {
278
- let acc = 0; // reduce pass, ascending z
279
- for (let z = 0; z < S; z++) {
280
- const p0 = z * ks, p1 = Math.min(k, p0 + ks);
281
- let s = 0; // one partial, in order
282
- for (let p = p0; p < p1; p++) {
283
- const a = transA ? A[p * m + row] : A[row * k + p];
284
- s = fma ? f32(s + a * Bm[p * n + col]) // single rounding
285
- : f32(s + f32(a * Bm[p * n + col]));
286
- }
287
- acc = f32(acc + s);
288
- }
289
- out[row * n + col] = acc;
290
- }
291
- return out;
292
- }
293
-
294
- // ---- live audits for the fused attention kernels ---------------------------
295
- // The attention kernels had exact INIT gates but nothing at live shapes —
296
- // the exact gap the GEMM audit exists to close, left open on the kernels with
297
- // the trickiest indexing (head-strided gather, scatter write-back). These
298
- // recompute individual output cells from the units, stratified like
299
- // auditTile: last/first token pair, last head, last channel first, then
300
- // random. Cost is hd (or T) multiply-adds per cell.
301
- function auditAttScores(qq, kq, qs, ks, d, got, L, nCells) {
302
- const { B, T, heads, hd } = d, C = heads * hd, mul = L.mul, raw = !!d.acc;
303
- const N = nCells || 8;
304
- const edges = [[B - 1, heads - 1, T - 1, T - 1], [0, 0, 0, 0],
305
- [0, heads - 1, T - 1, 0], [B - 1, 0, 0, T - 1]];
306
- for (let t = 0; t < N; t++) {
307
- let bi, h, ti, tj;
308
- if (t < edges.length) { bi = edges[t][0]; h = edges[t][1]; ti = edges[t][2]; tj = edges[t][3]; }
309
- else { bi = (Math.random() * B) | 0; h = (Math.random() * heads) | 0;
310
- ti = (Math.random() * T) | 0; tj = (Math.random() * T) | 0; }
311
- const bz = bi * heads + h;
312
- const qo = (bi * T + ti) * C + h * hd, ko = (bi * T + tj) * C + h * hd;
313
- let acc = 0;
314
- for (let p = 0; p < hd; p++) acc += mul[(qq[qo + p] & 0xFF) * 256 + (kq[ko + p] & 0xFF)];
315
- const v = raw ? acc : epi(acc, qs[(bi * T + ti) * heads + h], ks[(bi * T + tj) * heads + h]);
316
- const idx = (bz * T + ti) * T + tj;
317
- if (raw ? got[idx] !== v : bitDiff(got[idx], v))
318
- return `att.scores audit failed at [b${bi},h${h},${ti},${tj}] B${B}T${T}H${heads}d${hd}: kernel ${got[idx]} vs units ${v}`;
319
- }
320
- return null;
321
- }
322
- function auditAttCtx(aq, vq, as, vs, d, got, L, nCells) {
323
- const { B, T, heads, hd } = d, C = heads * hd, mul = L.mul, raw = !!d.acc;
324
- const N = nCells || 8;
325
- const edges = [[B - 1, heads - 1, T - 1, hd - 1], [0, 0, 0, 0],
326
- [0, heads - 1, T - 1, 0], [B - 1, 0, 0, hd - 1]];
327
- for (let t = 0; t < N; t++) {
328
- let bi, h, ti, j;
329
- if (t < edges.length) { bi = edges[t][0]; h = edges[t][1]; ti = edges[t][2]; j = edges[t][3]; }
330
- else { bi = (Math.random() * B) | 0; h = (Math.random() * heads) | 0;
331
- ti = (Math.random() * T) | 0; j = (Math.random() * hd) | 0; }
332
- const bz = bi * heads + h, ao = (bz * T + ti) * T;
333
- let acc = 0;
334
- for (let tj = 0; tj < T; tj++)
335
- acc += mul[(aq[ao + tj] & 0xFF) * 256 + (vq[(bi * T + tj) * C + h * hd + j] & 0xFF)];
336
- const v = raw ? acc : epi(acc, as[bz * T + ti], vs[(bi * heads + h) * hd + j]);
337
- const idx = (bi * T + ti) * C + h * hd + j;
338
- if (raw ? got[idx] !== v : bitDiff(got[idx], v))
339
- return `att.ctx audit failed at [b${bi},h${h},${ti},${j}] B${B}T${T}H${heads}d${hd}: kernel ${got[idx]} vs units ${v}`;
340
- }
341
- return null;
342
- }
343
-
344
- // block-scaled verified GEMM, float in → float out.
345
- // d = { m, k, n, batch=1, relu=false }; X is (batch·m)×k, W is batch×(k×n)
346
- // gpuBgemm (from webgpu.js) runs the batched kernel with the fused epilogue;
347
- // without it the CPU LUT mirror runs. Every product goes through the units.
348
- async function vgemmBlock(Xf, Wf, d, L, gpuBgemm, audit) {
349
- const { m, k, n } = d, batch = d.batch || 1;
350
- const x = quantizeRows(Xf, batch * m, k);
351
- let wq, ws;
352
- if (batch === 1) {
353
- const w = quantizeCols(Wf, k, n); wq = w.q; ws = w.s;
354
- } else {
355
- wq = new Int8Array(batch * k * n); ws = new Float32Array(batch * n);
356
- for (let bz = 0; bz < batch; bz++) {
357
- const w = quantizeCols(Wf.subarray(bz * k * n, (bz + 1) * k * n), k, n);
358
- wq.set(w.q, bz * k * n); ws.set(w.s, bz * n);
359
- }
360
- }
361
- if (gpuBgemm) {
362
- const out = await gpuBgemm(x.q, wq, x.s, ws, d);
363
- // continuous re-verification at LIVE shapes: the boot gate only ever saw
364
- // toy inputs, so sample a few real cells against the units as we go
365
- if (audit && audit.due()) {
366
- const bad = auditTile(x.q, wq, x.s, ws, d, out, L, audit.cells);
367
- if (bad) audit.fail(bad);
368
- }
369
- return out;
370
- }
371
- return bgemmJS(x.q, wq, x.s, ws, d, L);
372
- }
373
-
374
- // ---- gather-fused attention through the units (CUTLASS ex. 36/52) ----------
375
- // The kernels read q/k/v/ctx directly in their natural BT×C layout with
376
- // head-strided indexing — no JS gather copies, no kᵀ transpose, and the
377
- // context write scatters straight back into BT×C. Quantization stays
378
- // block-scaled: q/k/a per (token,head) row, v per (head,channel) column.
379
- // The (BT·heads)×hd row view of q/k IS the contiguous buffer, so
380
- // quantizeRows(q, BT·heads, hd) gives per-(token,head) scales for free.
381
- function quantizeHeadCols(v, B, T, heads, hd) { // per (batch,head,channel) column
382
- const C = heads * hd;
383
- const q = new Int8Array(B * T * C), s = new Float32Array(B * heads * hd);
384
- for (let bi = 0; bi < B; bi++)
385
- for (let h = 0; h < heads; h++)
386
- for (let j = 0; j < hd; j++) {
387
- let mx = 0;
388
- for (let ti = 0; ti < T; ti++) {
389
- const a = Math.abs(v[(bi * T + ti) * C + h * hd + j]);
390
- if (a > mx) mx = a;
391
- }
392
- const sc = Math.max(mx / 127, 1e-8);
393
- s[(bi * heads + h) * hd + j] = sc;
394
- for (let ti = 0; ti < T; ti++) {
395
- const idx = (bi * T + ti) * C + h * hd + j;
396
- const w = Math.round(v[idx] / sc);
397
- q[idx] = w < -128 ? -128 : w > 127 ? 127 : w;
398
- }
399
- }
400
- return { q, s };
401
- }
402
- // scores S[bz,ti,tj] = q_row(bi,ti,h) · k_row(bi,tj,h), every product via the LUT
403
- // d.acc=true returns the raw int32 accumulator (exact oracle for the kernel gate)
404
- function attScoresJS(qq, kq, qs, ks, d, L) {
405
- const { B, T, heads, hd } = d, C = heads * hd, mul = L.mul, raw = !!d.acc;
406
- const out = raw ? new Int32Array(B * heads * T * T) : new Float32Array(B * heads * T * T);
407
- for (let bi = 0; bi < B; bi++) for (let h = 0; h < heads; h++) {
408
- const bz = bi * heads + h;
409
- for (let ti = 0; ti < T; ti++) {
410
- const qo = (bi * T + ti) * C + h * hd, rscale = qs[(bi * T + ti) * heads + h];
411
- for (let tj = 0; tj < T; tj++) {
412
- const ko = (bi * T + tj) * C + h * hd;
413
- let acc = 0;
414
- for (let p = 0; p < hd; p++) acc += mul[(qq[qo + p] & 0xFF) * 256 + (kq[ko + p] & 0xFF)];
415
- out[(bz * T + ti) * T + tj] = raw ? acc : epi(acc, rscale, ks[(bi * T + tj) * heads + h]);
416
- }
417
- }
418
- }
419
- return out;
420
- }
421
- // ctx[(bi,ti),(h,j)] = Σ_tj a[bz,ti,tj]·v[(bi,tj),(h,j)] — scatter fused into BT×C
422
- function attCtxJS(aq, vq, as, vs, d, L) {
423
- const { B, T, heads, hd } = d, C = heads * hd, mul = L.mul, raw = !!d.acc;
424
- const out = raw ? new Int32Array(B * T * C) : new Float32Array(B * T * C);
425
- for (let bi = 0; bi < B; bi++) for (let h = 0; h < heads; h++) {
426
- const bz = bi * heads + h;
427
- for (let ti = 0; ti < T; ti++) {
428
- const ao = (bz * T + ti) * T, rscale = as[bz * T + ti];
429
- for (let j = 0; j < hd; j++) {
430
- let acc = 0;
431
- for (let tj = 0; tj < T; tj++)
432
- acc += mul[(aq[ao + tj] & 0xFF) * 256 + (vq[(bi * T + tj) * C + h * hd + j] & 0xFF)];
433
- out[(bi * T + ti) * C + h * hd + j] = raw ? acc : epi(acc, rscale, vs[(bi * heads + h) * hd + j]);
434
- }
435
- }
436
- }
437
- return out;
438
- }
439
-
440
- // one verified layer forward; returns float out (+ cache for STE backward).
441
- // Every product goes through the verified INT8 multiply (mul8 LUT) with exact
442
- // int32 accumulation — i.e. an emulated INT8 tensor-core GEMM — then dequant.
443
- async function linearFwd(X, W, m, k, n, L, useRelu, matmulInt8) {
444
- const xq = quantize(X), wq = quantize(W);
445
- const acc = await (matmulInt8 || lutMatmulJS)(xq.q, wq.q, m, k, n, L); // verified multiply
446
- const dq = xq.scale * wq.scale;
447
- const out = new Float32Array(m * n);
448
- const mask = useRelu ? new Uint8Array(m * n) : null;
449
- for (let i = 0; i < m * n; i++) {
450
- let v = acc[i] * dq;
451
- if (useRelu) { if (v > 0) mask[i] = 1; else v = 0; }
452
- out[i] = v;
453
- }
454
- return { out, mask };
455
- }
456
-
457
- // 2-layer MLP: X→H (relu) →dout. Forward through verified units, MSE loss.
458
- async function forward(X, y, W1, W2, D, L, matmulInt8) {
459
- const { n, din, h, dout } = D;
460
- const l1 = await linearFwd(X, W1, n, din, h, L, true, matmulInt8);
461
- const l2 = await linearFwd(l1.out, W2, n, h, dout, L, false, matmulInt8);
462
- const resid = new Float32Array(n * dout); let loss = 0;
463
- for (let i = 0; i < resid.length; i++) { const r = l2.out[i] - y[i]; resid[i] = r; loss += r * r; }
464
- loss /= resid.length;
465
- return { loss, resid, z1: l1.out, mask1: l1.mask };
466
- }
467
-
468
- // STE backward (verified matmul treated as float X@W). Returns flat [gW1, gW2].
469
- function backward(X, W1, W2, fwd, D) {
470
- const { n, din, h, dout } = D;
471
- const { resid, z1, mask1 } = fwd;
472
- const s = 2 / n;
473
- const dout_ = new Float32Array(resid.length);
474
- for (let i = 0; i < resid.length; i++) dout_[i] = resid[i] * s;
475
- const mm = TC.matmul, tr = TC.transpose;
476
- const gW2 = mm(tr(z1, n, h), dout_, h, n, dout); // z1ᵀ @ dout
477
- const dz1 = mm(dout_, tr(W2, h, dout), n, dout, h); // dout @ W2ᵀ
478
- for (let i = 0; i < dz1.length; i++) if (!mask1[i]) dz1[i] = 0; // relu grad
479
- const gW1 = mm(tr(X, n, din), dz1, din, n, h); // Xᵀ @ dz1
480
- const g = new Float32Array(gW1.length + gW2.length);
481
- g.set(gW1, 0); g.set(gW2, gW1.length);
482
- return g;
483
- }
484
-
485
- function splitApply(W1, W2, gAvg, lr) {
486
- for (let i = 0; i < W1.length; i++) W1[i] -= lr * gAvg[i];
487
- for (let j = 0; j < W2.length; j++) W2[j] -= lr * gAvg[W1.length + j];
488
- }
489
-
490
- const api = { quantize, quantize2, quantizeRows, quantizeCols, quantizeHeadCols, lutMatmulJS, lutMatmul3JS, lutMatmul3,
491
- bgemmJS, vgemmBlock, auditTile, epi, attScoresJS, attCtxJS, linearFwd, forward, backward, splitApply,
492
- rowAbsMax, scalesFromAbsMax, quantizeRowsInv, vmlpBlock, bitDiff,
493
- auditAttScores, auditAttCtx, fgemmMirror };
494
- if (typeof module !== "undefined" && module.exports) { TC = require("./traincore.js"); module.exports = api; }
495
- else { TC = root.TrainCore; root.Verified = api; }
496
- })(typeof self !== "undefined" ? self : this);
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ // Verified INT8 compute — the emulated GPU logic, in the browser.
2
+ // A layer's forward runs THROUGH the units: quantize -> LUT multiply -> requant
3
+ // -> optional ReLU -> dequant. Backward is a straight-through estimator (the
4
+ // integer path has no gradient), so ordinary float weights still learn.
5
+ // Same units as the Python/Docker DaisyChain; here they're lookup tables.
6
+ (function (root) {
7
+ "use strict";
8
+
9
+ let TC; // TrainCore (matmul/transpose) — resolved per environment at the end
10
+
11
+ // A float -> int8 conversion has no answer for NaN/Inf, and the quantizers below give a
12
+ // silent wrong one: an Inf makes the |max| scale Infinity and EVERY value quantizes to 0
13
+ // ([0.5, Inf, 3] -> [0, 0, 0]); a NaN is skipped by `a > mx` and then stored as 0 by the
14
+ // Int8Array. The Python trainer refuses the same input (qat.py); so does this. One pass,
15
+ // no effect on finite inputs, so builds with and without it still co-train.
16
+ function assertFinite(X, where) {
17
+ for (let i = 0; i < X.length; i++) {
18
+ if (!Number.isFinite(X[i])) {
19
+ let bad = 0;
20
+ for (let j = i; j < X.length; j++) if (!Number.isFinite(X[j])) bad++;
21
+ const e = new Error("non-finite input to " + where + " (" + bad + " of " + X.length +
22
+ " values): the int8 quantize would turn it into 0");
23
+ e.name = "NonFiniteError";
24
+ throw e;
25
+ }
26
+ }
27
+ }
28
+
29
+ function quantize(X) {
30
+ assertFinite(X, "quantize");
31
+ let mx = 0; for (let i = 0; i < X.length; i++) { const a = Math.abs(X[i]); if (a > mx) mx = a; }
32
+ const scale = Math.max(mx / 127, 1e-8);
33
+ const q = new Int8Array(X.length);
34
+ for (let i = 0; i < X.length; i++) { let v = Math.round(X[i] / scale); q[i] = v < -128 ? -128 : v > 127 ? 127 : v; }
35
+ return { q, scale };
36
+ }
37
+
38
+ // int8 matmul via the verified multiply LUT: acc(m×n) = sum_k mulLUT[Xq,Wq]
39
+ function lutMatmulJS(Xq, Wq, m, k, n, L) {
40
+ const C = new Int32Array(m * n), mul = L.mul;
41
+ for (let i = 0; i < m; i++) {
42
+ for (let p = 0; p < k; p++) {
43
+ const au = (Xq[i * k + p] & 0xFF) * 256, wo = p * n, co = i * n;
44
+ for (let j = 0; j < n; j++) C[co + j] += mul[au + (Wq[wo + j] & 0xFF)];
45
+ }
46
+ }
47
+ return C;
48
+ }
49
+
50
+ // ---- 3xINT8 fast-accurate GEMM --------------------------------------------
51
+ // The CUTLASS example-27 "3xTF32" scheme, ported to the verified units:
52
+ // split each float into a coarse int8 part plus an int8-quantized residual,
53
+ // run three EXACT LUT GEMMs (hi·hi, hi·lo, lo·hi), drop the negligible
54
+ // lo·lo, and recombine. Same big/small decomposition NVIDIA uses to recover
55
+ // near-fp32 accuracy from TF32 tensor cores — here it recovers ~14-bit
56
+ // accuracy from the 8-bit units, at 3× the unit ops. Every product still
57
+ // goes through the verified mul8 LUT.
58
+ function quantize2(X) {
59
+ const hi = quantize(X);
60
+ const r = new Float32Array(X.length);
61
+ for (let i = 0; i < X.length; i++) r[i] = X[i] - hi.q[i] * hi.scale;
62
+ const lo = quantize(r);
63
+ return { hi, lo };
64
+ }
65
+ function combine3(hh, hl, lh, x, w, len) {
66
+ const out = new Float32Array(len);
67
+ const shh = x.hi.scale * w.hi.scale, shl = x.hi.scale * w.lo.scale, slh = x.lo.scale * w.hi.scale;
68
+ for (let i = 0; i < len; i++) out[i] = hh[i] * shh + hl[i] * shl + lh[i] * slh;
69
+ return out;
70
+ }
71
+ function lutMatmul3JS(Xf, Wf, m, k, n, L) { // sync, CPU LUT path
72
+ const x = quantize2(Xf), w = quantize2(Wf);
73
+ return combine3(lutMatmulJS(x.hi.q, w.hi.q, m, k, n, L),
74
+ lutMatmulJS(x.hi.q, w.lo.q, m, k, n, L),
75
+ lutMatmulJS(x.lo.q, w.hi.q, m, k, n, L), x, w, m * n);
76
+ }
77
+ async function lutMatmul3(Xf, Wf, m, k, n, L, matmulInt8) { // any backend
78
+ const x = quantize2(Xf), w = quantize2(Wf);
79
+ const mm = matmulInt8 || lutMatmulJS;
80
+ const [hh, hl, lh] = await Promise.all([
81
+ mm(x.hi.q, w.hi.q, m, k, n, L),
82
+ mm(x.hi.q, w.lo.q, m, k, n, L),
83
+ mm(x.lo.q, w.hi.q, m, k, n, L),
84
+ ]);
85
+ return combine3(hh, hl, lh, x, w, m * n);
86
+ }
87
+
88
+ // ---- block-scaled verified GEMM (CUTLASS ex. 67/81 blockwise scaling) ------
89
+ // Per-ROW scales for the activations and per-COLUMN scales for the weights:
90
+ // the integer math through the mul8 LUT is completely unchanged — only the
91
+ // dequant uses rs[row]·cs[col] instead of one tensor-wide product, so a single
92
+ // outlier no longer crushes the quantization resolution of every other
93
+ // row/column. One LUT pass at this granularity beats the per-tensor 3-pass.
94
+ function quantizeRows(X, rows, cols) {
95
+ assertFinite(X, "quantizeRows");
96
+ const q = new Int8Array(rows * cols), s = new Float32Array(rows);
97
+ for (let r = 0; r < rows; r++) {
98
+ let mx = 0;
99
+ for (let c = 0; c < cols; c++) { const a = Math.abs(X[r * cols + c]); if (a > mx) mx = a; }
100
+ const sc = Math.max(mx / 127, 1e-8); s[r] = sc;
101
+ for (let c = 0; c < cols; c++) {
102
+ const v = Math.round(X[r * cols + c] / sc);
103
+ q[r * cols + c] = v < -128 ? -128 : v > 127 ? 127 : v;
104
+ }
105
+ }
106
+ return { q, s };
107
+ }
108
+ function quantizeCols(W, rows, cols) {
109
+ assertFinite(W, "quantizeCols");
110
+ const q = new Int8Array(rows * cols), s = new Float32Array(cols);
111
+ for (let c = 0; c < cols; c++) {
112
+ let mx = 0;
113
+ for (let r = 0; r < rows; r++) { const a = Math.abs(W[r * cols + c]); if (a > mx) mx = a; }
114
+ s[c] = Math.max(mx / 127, 1e-8);
115
+ }
116
+ for (let r = 0; r < rows; r++)
117
+ for (let c = 0; c < cols; c++) {
118
+ const v = Math.round(W[r * cols + c] / s[c]);
119
+ q[r * cols + c] = v < -128 ? -128 : v > 127 ? 127 : v;
120
+ }
121
+ return { q, s };
122
+ }
123
+
124
+ // ---- B2B MLP chain (CUTLASS ex. 13 two-GEMM fusion + ex. 23 epilogue
125
+ // reduction), respecced for cross-device exactness ---------------------------
126
+ // The MLP is the one back-to-back GEMM pair with no layernorm/softmax between
127
+ // (ReLU is already fused in the epilogue), so the intermediate h1 can be
128
+ // quantized ON the GPU and fed straight to the second GEMM. Two rules make
129
+ // that fleet-safe:
130
+ // 1. The per-row |max| (ex. 23) uses only comparisons — exact on any
131
+ // hardware, order-independent — and comes back to JS as ~1KB.
132
+ // 2. Scale DERIVATION (two divisions) stays in JS f64, which IEEE requires
133
+ // to be exactly rounded and is therefore identical on every device.
134
+ // WGSL division is only 2.5 ULP — a fork waiting to happen — but WGSL
135
+ // multiply/add are correctly rounded and floor/clamp are exact. So the
136
+ // quantize step is respecced from round(x / scale) to
137
+ // floor(f32(x * invScale) + 0.5) — floor(x+0.5) IS Math.round's tie
138
+ // rule — and the fround-stepped mirror below is bit-identical to the
139
+ // GPU kernel, which is exact-gated against it at init.
140
+ // NOTE this changes which int8 a value on a rounding boundary lands on
141
+ // (≤1 step) vs quantizeRows, so old and new builds cannot co-train — the
142
+ // divergence guard stops such mixed groups by design.
143
+ function rowAbsMax(X, rows, cols) {
144
+ assertFinite(X, "rowAbsMax");
145
+ const mx = new Float32Array(rows);
146
+ for (let r = 0; r < rows; r++) {
147
+ let m = 0;
148
+ for (let c = 0; c < cols; c++) { const a = Math.abs(X[r * cols + c]); if (a > m) m = a; }
149
+ mx[r] = m;
150
+ }
151
+ return mx;
152
+ }
153
+ function scalesFromAbsMax(mx) { // f64 divisions: exactly rounded, device-identical
154
+ const scale = new Float32Array(mx.length), inv = new Float32Array(mx.length);
155
+ for (let i = 0; i < mx.length; i++) {
156
+ scale[i] = Math.max(mx[i] / 127, 1e-8);
157
+ inv[i] = 1 / scale[i]; // recip of the STORED f32 scale
158
+ }
159
+ return { scale, inv };
160
+ }
161
+ function quantizeRowsInv(X, rows, cols, inv) { // bit-exact mirror of the GPU quantize kernel
162
+ const q = new Int8Array(rows * cols);
163
+ for (let r = 0; r < rows; r++) {
164
+ const iv = inv[r];
165
+ for (let c = 0; c < cols; c++) {
166
+ const n = Math.floor(f32(f32(X[r * cols + c] * iv) + 0.5));
167
+ q[r * cols + c] = n < -128 ? -128 : n > 127 ? 127 : n;
168
+ }
169
+ }
170
+ return q;
171
+ }
172
+ // the chained MLP: X @ W1 -> ReLU (fused) -> absmax -> quantize -> @ W2.
173
+ // d = { m, k, h, n }; gpuMlp (from webgpu.js) runs both GEMMs + the
174
+ // on-GPU quantize with one tiny absmax readback between; without it the CPU
175
+ // mirror chain runs — SAME math, so mixed GPU/CPU fleets stay bit-identical.
176
+ async function vmlpBlock(Xf, W1f, W2f, d, L, gpuMlp, audit) {
177
+ const x = quantizeRows(Xf, d.m, d.k);
178
+ const w1 = quantizeCols(W1f, d.k, d.h);
179
+ const w2 = quantizeCols(W2f, d.h, d.n);
180
+ if (gpuMlp) {
181
+ const r = await gpuMlp(x.q, w1.q, w2.q, x.s, w1.s, w2.s, d);
182
+ if (audit && audit.due()) {
183
+ // audit BOTH live GEMMs: gemm1 against the units directly; gemm2 by
184
+ // reconstructing its exact operand through the proven quantize mirror
185
+ const bad1 = auditTile(x.q, w1.q, x.s, w1.s, { m: d.m, k: d.k, n: d.h, relu: true }, r.h1, L, audit.cells);
186
+ if (bad1) { audit.fail("mlp gemm1: " + bad1); return r; }
187
+ const sc = scalesFromAbsMax(rowAbsMax(r.h1, d.m, d.h));
188
+ const hq = quantizeRowsInv(r.h1, d.m, d.h, sc.inv);
189
+ const bad2 = auditTile(hq, w2.q, sc.scale, w2.s, { m: d.m, k: d.h, n: d.n }, r.out, L, audit.cells);
190
+ if (bad2) audit.fail("mlp gemm2: " + bad2);
191
+ }
192
+ return r;
193
+ }
194
+ const h1 = bgemmJS(x.q, w1.q, x.s, w1.s, { m: d.m, k: d.k, n: d.h, batch: 1, relu: true }, L);
195
+ const sc = scalesFromAbsMax(rowAbsMax(h1, d.m, d.h));
196
+ const hq = quantizeRowsInv(h1, d.m, d.h, sc.inv);
197
+ const out = bgemmJS(hq, w2.q, sc.scale, w2.s, { m: d.m, k: d.h, n: d.n, batch: 1 }, L);
198
+ return { h1, out };
199
+ }
200
+
201
+ // ---- epilogue mirror -------------------------------------------------------
202
+ // BIT-EXACT mirror of the WGSL epilogue `f32(s) * a * b`. WGSL rounds to f32
203
+ // after the int->float conversion and after EACH multiply; plain JS would do
204
+ // the whole chain in f64 and round once, which differs in the last ulp. That
205
+ // last-ulp gap is what used to force a tolerance into the kernel gates —
206
+ // mirroring the rounding exactly is what lets the gates compare with `!==`.
207
+ const f32 = Math.fround;
208
+ function epi(s, a, b) { return f32(f32(f32(s) * a) * b); }
209
+
210
+ // f32 equality at the BIT level: `!==` says -0 === 0, but replicas are
211
+ // compared by hashing raw bytes, so an audit that can't see the sign of zero
212
+ // could pass a device that later forks the fleet. (Real ISAs have non-IEEE
213
+ // modes that flush -0 to +0 — e.g. RDNA2 output modifiers / legacy muls.)
214
+ const _fb = new Float32Array(1), _ub = new Uint32Array(_fb.buffer);
215
+ function bitDiff(a, b) { _fb[0] = a; const u = _ub[0]; _fb[0] = b; return u !== _ub[0]; }
216
+
217
+ // CPU mirror of the fused GPU kernel: batched int8 GEMM through the LUT with
218
+ // the epilogue (block dequant + optional ReLU) applied before returning —
219
+ // exactly what the WGSL kernel does on-device. d.acc=true returns the raw
220
+ // int32 accumulator instead (the exact oracle the fused kernel normally hides).
221
+ function bgemmJS(Xq, Wq, rs, cs, d, L) {
222
+ const { m, k, n } = d, batch = d.batch || 1, relu = !!d.relu, mul = L.mul;
223
+ const raw = !!d.acc;
224
+ const out = raw ? new Int32Array(batch * m * n) : new Float32Array(batch * m * n);
225
+ const acc = new Int32Array(n);
226
+ for (let bz = 0; bz < batch; bz++) {
227
+ const xo = bz * m * k, wo = bz * k * n, oo = bz * m * n, co = bz * n;
228
+ for (let i = 0; i < m; i++) {
229
+ acc.fill(0);
230
+ const xrow = xo + i * k;
231
+ for (let p = 0; p < k; p++) {
232
+ const au = (Xq[xrow + p] & 0xFF) * 256, wrow = wo + p * n;
233
+ for (let j = 0; j < n; j++) acc[j] += mul[au + (Wq[wrow + j] & 0xFF)];
234
+ }
235
+ const orow = oo + i * n;
236
+ if (raw) { for (let j = 0; j < n; j++) out[orow + j] = acc[j]; continue; }
237
+ const rscale = rs[bz * m + i];
238
+ for (let j = 0; j < n; j++) {
239
+ const v = epi(acc[j], rscale, cs[co + j]);
240
+ out[orow + j] = relu && v < 0 ? 0 : v;
241
+ }
242
+ }
243
+ }
244
+ return out;
245
+ }
246
+
247
+ // Recompute a handful of RANDOM output cells of a live GEMM through the LUT
248
+ // mirror and compare against what the kernel produced. Sampling cells instead
249
+ // of whole matrices makes this cheap enough to run continuously, at the real
250
+ // shapes training uses — not once at boot on toy inputs.
251
+ // STRATIFIED sampling. Uniformly random cells are the wrong instrument for
252
+ // the bugs that actually occur here: a bounds-guard off-by-one or a pack-tail
253
+ // padding bug lives on the LAST row/column, and uniform sampling finds that
254
+ // with probability ~1/n per cell — at the 16512-wide logits GEMM, never. So
255
+ // the first cells are the structurally dangerous ones (corners, last row,
256
+ // last column, last batch) chosen deterministically, and the remainder are
257
+ // random interior cells that catch diffuse bugs. Same principle as poisoning
258
+ // the buffer pool: construct the dangerous case, don't wait to land on it.
259
+ function auditTile(Xq, Wq, rs, cs, d, got, L, nCells) {
260
+ const { m, k, n } = d, batch = d.batch || 1, relu = !!d.relu, mul = L.mul;
261
+ const N = nCells || 8;
262
+ const edges = [[0, m - 1, n - 1], [0, 0, n - 1], [0, m - 1, 0], [0, 0, 0],
263
+ [batch - 1, m - 1, n - 1], [batch - 1, 0, 0]];
264
+ for (let t = 0; t < N; t++) {
265
+ let bz, i, j;
266
+ if (t < edges.length) { bz = edges[t][0]; i = edges[t][1]; j = edges[t][2]; }
267
+ else { bz = (Math.random() * batch) | 0; i = (Math.random() * m) | 0; j = (Math.random() * n) | 0; }
268
+ let acc = 0;
269
+ const xrow = bz * m * k + i * k, wo = bz * k * n;
270
+ for (let p = 0; p < k; p++) acc += mul[(Xq[xrow + p] & 0xFF) * 256 + (Wq[wo + p * n + j] & 0xFF)];
271
+ let v = epi(acc, rs[bz * m + i], cs[bz * n + j]);
272
+ if (relu && v < 0) v = 0;
273
+ const idx = (bz * m + i) * n + j;
274
+ if (bitDiff(got[idx], v))
275
+ return `GEMM audit failed at [b${bz},${i},${j}] shape ${m}x${k}x${n}: kernel ${Object.is(got[idx], -0) ? "-0" : got[idx]} vs units ${Object.is(v, -0) ? "-0" : v}`;
276
+ }
277
+ return null;
278
+ }
279
+
280
+ // ---- exact mirror of the split-K f32 GEMM ----------------------------------
281
+ // The f32 backward GEMM was the last kernel gated by a TOLERANCE (allclose at
282
+ // 1e-3) — and this project's own gate mutation test shows allclose waving
283
+ // through real bugs. The reason was real though: split-K accumulates in a
284
+ // different ORDER than a naive reference, so bit-equality against the naive
285
+ // one is impossible. The fix is the same as the epilogue mirror: reproduce
286
+ // the kernel's order exactly, then compare with `!==`.
287
+ // partials: for z in 0..S-1, sum p in [z*ks, min(k,(z+1)*ks)) in order
288
+ // reduce: sum the S partials in ascending z
289
+ // `fma` selects the rounding schedule for `s + a*b`: WGSL PERMITS a compiler
290
+ // to contract that into a fused multiply-add (one rounding) instead of two.
291
+ // Which one the device does is a fact about the device, so the gate tries
292
+ // both and reports which matches rather than assuming.
293
+ function fgemmMirror(A, Bm, d, fma) {
294
+ const { m, k, n } = d, transA = !!d.transA;
295
+ const S = k > 4096 ? Math.min(16, Math.ceil(k / 2048)) : 1;
296
+ const ks = Math.ceil(k / S);
297
+ const out = new Float32Array(m * n);
298
+ for (let row = 0; row < m; row++)
299
+ for (let col = 0; col < n; col++) {
300
+ let acc = 0; // reduce pass, ascending z
301
+ for (let z = 0; z < S; z++) {
302
+ const p0 = z * ks, p1 = Math.min(k, p0 + ks);
303
+ let s = 0; // one partial, in order
304
+ for (let p = p0; p < p1; p++) {
305
+ const a = transA ? A[p * m + row] : A[row * k + p];
306
+ s = fma ? f32(s + a * Bm[p * n + col]) // single rounding
307
+ : f32(s + f32(a * Bm[p * n + col]));
308
+ }
309
+ acc = f32(acc + s);
310
+ }
311
+ out[row * n + col] = acc;
312
+ }
313
+ return out;
314
+ }
315
+
316
+ // ---- live audits for the fused attention kernels ---------------------------
317
+ // The attention kernels had exact INIT gates but nothing at live shapes —
318
+ // the exact gap the GEMM audit exists to close, left open on the kernels with
319
+ // the trickiest indexing (head-strided gather, scatter write-back). These
320
+ // recompute individual output cells from the units, stratified like
321
+ // auditTile: last/first token pair, last head, last channel first, then
322
+ // random. Cost is hd (or T) multiply-adds per cell.
323
+ function auditAttScores(qq, kq, qs, ks, d, got, L, nCells) {
324
+ const { B, T, heads, hd } = d, C = heads * hd, mul = L.mul, raw = !!d.acc;
325
+ const N = nCells || 8;
326
+ const edges = [[B - 1, heads - 1, T - 1, T - 1], [0, 0, 0, 0],
327
+ [0, heads - 1, T - 1, 0], [B - 1, 0, 0, T - 1]];
328
+ for (let t = 0; t < N; t++) {
329
+ let bi, h, ti, tj;
330
+ if (t < edges.length) { bi = edges[t][0]; h = edges[t][1]; ti = edges[t][2]; tj = edges[t][3]; }
331
+ else { bi = (Math.random() * B) | 0; h = (Math.random() * heads) | 0;
332
+ ti = (Math.random() * T) | 0; tj = (Math.random() * T) | 0; }
333
+ const bz = bi * heads + h;
334
+ const qo = (bi * T + ti) * C + h * hd, ko = (bi * T + tj) * C + h * hd;
335
+ let acc = 0;
336
+ for (let p = 0; p < hd; p++) acc += mul[(qq[qo + p] & 0xFF) * 256 + (kq[ko + p] & 0xFF)];
337
+ const v = raw ? acc : epi(acc, qs[(bi * T + ti) * heads + h], ks[(bi * T + tj) * heads + h]);
338
+ const idx = (bz * T + ti) * T + tj;
339
+ if (raw ? got[idx] !== v : bitDiff(got[idx], v))
340
+ return `att.scores audit failed at [b${bi},h${h},${ti},${tj}] B${B}T${T}H${heads}d${hd}: kernel ${got[idx]} vs units ${v}`;
341
+ }
342
+ return null;
343
+ }
344
+ function auditAttCtx(aq, vq, as, vs, d, got, L, nCells) {
345
+ const { B, T, heads, hd } = d, C = heads * hd, mul = L.mul, raw = !!d.acc;
346
+ const N = nCells || 8;
347
+ const edges = [[B - 1, heads - 1, T - 1, hd - 1], [0, 0, 0, 0],
348
+ [0, heads - 1, T - 1, 0], [B - 1, 0, 0, hd - 1]];
349
+ for (let t = 0; t < N; t++) {
350
+ let bi, h, ti, j;
351
+ if (t < edges.length) { bi = edges[t][0]; h = edges[t][1]; ti = edges[t][2]; j = edges[t][3]; }
352
+ else { bi = (Math.random() * B) | 0; h = (Math.random() * heads) | 0;
353
+ ti = (Math.random() * T) | 0; j = (Math.random() * hd) | 0; }
354
+ const bz = bi * heads + h, ao = (bz * T + ti) * T;
355
+ let acc = 0;
356
+ for (let tj = 0; tj < T; tj++)
357
+ acc += mul[(aq[ao + tj] & 0xFF) * 256 + (vq[(bi * T + tj) * C + h * hd + j] & 0xFF)];
358
+ const v = raw ? acc : epi(acc, as[bz * T + ti], vs[(bi * heads + h) * hd + j]);
359
+ const idx = (bi * T + ti) * C + h * hd + j;
360
+ if (raw ? got[idx] !== v : bitDiff(got[idx], v))
361
+ return `att.ctx audit failed at [b${bi},h${h},${ti},${j}] B${B}T${T}H${heads}d${hd}: kernel ${got[idx]} vs units ${v}`;
362
+ }
363
+ return null;
364
+ }
365
+
366
+ // block-scaled verified GEMM, float in → float out.
367
+ // d = { m, k, n, batch=1, relu=false }; X is (batch·m)×k, W is batch×(k×n)
368
+ // gpuBgemm (from webgpu.js) runs the batched kernel with the fused epilogue;
369
+ // without it the CPU LUT mirror runs. Every product goes through the units.
370
+ async function vgemmBlock(Xf, Wf, d, L, gpuBgemm, audit) {
371
+ const { m, k, n } = d, batch = d.batch || 1;
372
+ const x = quantizeRows(Xf, batch * m, k);
373
+ let wq, ws;
374
+ if (batch === 1) {
375
+ const w = quantizeCols(Wf, k, n); wq = w.q; ws = w.s;
376
+ } else {
377
+ wq = new Int8Array(batch * k * n); ws = new Float32Array(batch * n);
378
+ for (let bz = 0; bz < batch; bz++) {
379
+ const w = quantizeCols(Wf.subarray(bz * k * n, (bz + 1) * k * n), k, n);
380
+ wq.set(w.q, bz * k * n); ws.set(w.s, bz * n);
381
+ }
382
+ }
383
+ if (gpuBgemm) {
384
+ const out = await gpuBgemm(x.q, wq, x.s, ws, d);
385
+ // continuous re-verification at LIVE shapes: the boot gate only ever saw
386
+ // toy inputs, so sample a few real cells against the units as we go
387
+ if (audit && audit.due()) {
388
+ const bad = auditTile(x.q, wq, x.s, ws, d, out, L, audit.cells);
389
+ if (bad) audit.fail(bad);
390
+ }
391
+ return out;
392
+ }
393
+ return bgemmJS(x.q, wq, x.s, ws, d, L);
394
+ }
395
+
396
+ // ---- gather-fused attention through the units (CUTLASS ex. 36/52) ----------
397
+ // The kernels read q/k/v/ctx directly in their natural BT×C layout with
398
+ // head-strided indexing — no JS gather copies, no kᵀ transpose, and the
399
+ // context write scatters straight back into BT×C. Quantization stays
400
+ // block-scaled: q/k/a per (token,head) row, v per (head,channel) column.
401
+ // The (BT·heads)×hd row view of q/k IS the contiguous buffer, so
402
+ // quantizeRows(q, BT·heads, hd) gives per-(token,head) scales for free.
403
+ function quantizeHeadCols(v, B, T, heads, hd) { // per (batch,head,channel) column
404
+ assertFinite(v, "quantizeHeadCols");
405
+ const C = heads * hd;
406
+ const q = new Int8Array(B * T * C), s = new Float32Array(B * heads * hd);
407
+ for (let bi = 0; bi < B; bi++)
408
+ for (let h = 0; h < heads; h++)
409
+ for (let j = 0; j < hd; j++) {
410
+ let mx = 0;
411
+ for (let ti = 0; ti < T; ti++) {
412
+ const a = Math.abs(v[(bi * T + ti) * C + h * hd + j]);
413
+ if (a > mx) mx = a;
414
+ }
415
+ const sc = Math.max(mx / 127, 1e-8);
416
+ s[(bi * heads + h) * hd + j] = sc;
417
+ for (let ti = 0; ti < T; ti++) {
418
+ const idx = (bi * T + ti) * C + h * hd + j;
419
+ const w = Math.round(v[idx] / sc);
420
+ q[idx] = w < -128 ? -128 : w > 127 ? 127 : w;
421
+ }
422
+ }
423
+ return { q, s };
424
+ }
425
+ // scores S[bz,ti,tj] = q_row(bi,ti,h) · k_row(bi,tj,h), every product via the LUT
426
+ // d.acc=true returns the raw int32 accumulator (exact oracle for the kernel gate)
427
+ function attScoresJS(qq, kq, qs, ks, d, L) {
428
+ const { B, T, heads, hd } = d, C = heads * hd, mul = L.mul, raw = !!d.acc;
429
+ const out = raw ? new Int32Array(B * heads * T * T) : new Float32Array(B * heads * T * T);
430
+ for (let bi = 0; bi < B; bi++) for (let h = 0; h < heads; h++) {
431
+ const bz = bi * heads + h;
432
+ for (let ti = 0; ti < T; ti++) {
433
+ const qo = (bi * T + ti) * C + h * hd, rscale = qs[(bi * T + ti) * heads + h];
434
+ for (let tj = 0; tj < T; tj++) {
435
+ const ko = (bi * T + tj) * C + h * hd;
436
+ let acc = 0;
437
+ for (let p = 0; p < hd; p++) acc += mul[(qq[qo + p] & 0xFF) * 256 + (kq[ko + p] & 0xFF)];
438
+ out[(bz * T + ti) * T + tj] = raw ? acc : epi(acc, rscale, ks[(bi * T + tj) * heads + h]);
439
+ }
440
+ }
441
+ }
442
+ return out;
443
+ }
444
+ // ctx[(bi,ti),(h,j)] = Σ_tj a[bz,ti,tj]·v[(bi,tj),(h,j)] — scatter fused into BT×C
445
+ function attCtxJS(aq, vq, as, vs, d, L) {
446
+ const { B, T, heads, hd } = d, C = heads * hd, mul = L.mul, raw = !!d.acc;
447
+ const out = raw ? new Int32Array(B * T * C) : new Float32Array(B * T * C);
448
+ for (let bi = 0; bi < B; bi++) for (let h = 0; h < heads; h++) {
449
+ const bz = bi * heads + h;
450
+ for (let ti = 0; ti < T; ti++) {
451
+ const ao = (bz * T + ti) * T, rscale = as[bz * T + ti];
452
+ for (let j = 0; j < hd; j++) {
453
+ let acc = 0;
454
+ for (let tj = 0; tj < T; tj++)
455
+ acc += mul[(aq[ao + tj] & 0xFF) * 256 + (vq[(bi * T + tj) * C + h * hd + j] & 0xFF)];
456
+ out[(bi * T + ti) * C + h * hd + j] = raw ? acc : epi(acc, rscale, vs[(bi * heads + h) * hd + j]);
457
+ }
458
+ }
459
+ }
460
+ return out;
461
+ }
462
+
463
+ // one verified layer forward; returns float out (+ cache for STE backward).
464
+ // Every product goes through the verified INT8 multiply (mul8 LUT) with exact
465
+ // int32 accumulation — i.e. an emulated INT8 tensor-core GEMM — then dequant.
466
+ async function linearFwd(X, W, m, k, n, L, useRelu, matmulInt8) {
467
+ const xq = quantize(X), wq = quantize(W);
468
+ const acc = await (matmulInt8 || lutMatmulJS)(xq.q, wq.q, m, k, n, L); // verified multiply
469
+ const dq = xq.scale * wq.scale;
470
+ const out = new Float32Array(m * n);
471
+ const mask = useRelu ? new Uint8Array(m * n) : null;
472
+ for (let i = 0; i < m * n; i++) {
473
+ let v = acc[i] * dq;
474
+ if (useRelu) { if (v > 0) mask[i] = 1; else v = 0; }
475
+ out[i] = v;
476
+ }
477
+ return { out, mask };
478
+ }
479
+
480
+ // 2-layer MLP: X→H (relu) →dout. Forward through verified units, MSE loss.
481
+ async function forward(X, y, W1, W2, D, L, matmulInt8) {
482
+ const { n, din, h, dout } = D;
483
+ const l1 = await linearFwd(X, W1, n, din, h, L, true, matmulInt8);
484
+ const l2 = await linearFwd(l1.out, W2, n, h, dout, L, false, matmulInt8);
485
+ const resid = new Float32Array(n * dout); let loss = 0;
486
+ for (let i = 0; i < resid.length; i++) { const r = l2.out[i] - y[i]; resid[i] = r; loss += r * r; }
487
+ loss /= resid.length;
488
+ return { loss, resid, z1: l1.out, mask1: l1.mask };
489
+ }
490
+
491
+ // STE backward (verified matmul treated as float X@W). Returns flat [gW1, gW2].
492
+ function backward(X, W1, W2, fwd, D) {
493
+ const { n, din, h, dout } = D;
494
+ const { resid, z1, mask1 } = fwd;
495
+ const s = 2 / n;
496
+ const dout_ = new Float32Array(resid.length);
497
+ for (let i = 0; i < resid.length; i++) dout_[i] = resid[i] * s;
498
+ const mm = TC.matmul, tr = TC.transpose;
499
+ const gW2 = mm(tr(z1, n, h), dout_, h, n, dout); // z1ᵀ @ dout
500
+ const dz1 = mm(dout_, tr(W2, h, dout), n, dout, h); // dout @ W2ᵀ
501
+ for (let i = 0; i < dz1.length; i++) if (!mask1[i]) dz1[i] = 0; // relu grad
502
+ const gW1 = mm(tr(X, n, din), dz1, din, n, h); // Xᵀ @ dz1
503
+ const g = new Float32Array(gW1.length + gW2.length);
504
+ g.set(gW1, 0); g.set(gW2, gW1.length);
505
+ return g;
506
+ }
507
+
508
+ function splitApply(W1, W2, gAvg, lr) {
509
+ for (let i = 0; i < W1.length; i++) W1[i] -= lr * gAvg[i];
510
+ for (let j = 0; j < W2.length; j++) W2[j] -= lr * gAvg[W1.length + j];
511
+ }
512
+
513
+ const api = { assertFinite, quantize, quantize2, quantizeRows, quantizeCols, quantizeHeadCols, lutMatmulJS, lutMatmul3JS, lutMatmul3,
514
+ bgemmJS, vgemmBlock, auditTile, epi, attScoresJS, attCtxJS, linearFwd, forward, backward, splitApply,
515
+ rowAbsMax, scalesFromAbsMax, quantizeRowsInv, vmlpBlock, bitDiff,
516
+ auditAttScores, auditAttCtx, fgemmMirror };
517
+ if (typeof module !== "undefined" && module.exports) { TC = require("./traincore.js"); module.exports = api; }
518
+ else { TC = root.TrainCore; root.Verified = api; }
519
+ })(typeof self !== "undefined" ? self : this);
web/test_nonfinite.js ADDED
@@ -0,0 +1,25 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ // A float -> int8 quantize has no answer for NaN/Inf. Before the guard, an Inf made the |max|
2
+ // scale Infinity and every value quantized to 0 ([0.5, Inf, 3] -> [0, 0, 0]), and a NaN was
3
+ // skipped by the |max| scan and then stored as 0 by the Int8Array -- training continued on
4
+ // zeros with no error. Every quantizer must now refuse, and finite inputs must be unchanged.
5
+ "use strict";
6
+ const V = require("./public/verified_core.js");
7
+
8
+ let fails = 0;
9
+ function ck(name, cond) { console.log((cond ? " ok " : " FAIL ") + name); if (!cond) fails++; }
10
+ function throwsNonFinite(fn) { try { fn(); return false; } catch (e) { return e.name === "NonFiniteError"; } }
11
+
12
+ const bad = [Float32Array.from([0.5, Infinity, 3, -1]), Float32Array.from([0.5, NaN, 3, -1]),
13
+ Float32Array.from([-Infinity, 1, 2, 3])];
14
+ for (const X of bad) {
15
+ const tag = Array.from(X).join(",");
16
+ ck("quantize refuses [" + tag + "]", throwsNonFinite(() => V.quantize(X)));
17
+ ck("quantizeRows refuses [" + tag + "]", throwsNonFinite(() => V.quantizeRows(X, 2, 2)));
18
+ ck("quantizeCols refuses [" + tag + "]", throwsNonFinite(() => V.quantizeCols(X, 2, 2)));
19
+ ck("rowAbsMax refuses [" + tag + "]", throwsNonFinite(() => V.rowAbsMax(X, 2, 2)));
20
+ }
21
+ const q = V.quantize(Float32Array.from([0.5, -1, 3]));
22
+ ck("finite input quantizes as before: [21, -42, 127]", Array.from(q.q).join() === "21,-42,127");
23
+
24
+ if (fails) { console.log("NONFINITE TEST FAILED: " + fails); process.exit(1); }
25
+ console.log("NONFINITE TEST PASSED — every quantizer refuses NaN/Inf instead of zeroing.");