Refuse NaN/Inf at every int8 quantize; requant rounds half-up (retrained); bounded units; README for September
Browse filesFour 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 +58 -4
- daisychain/dashboard/agent.py +34 -1
- daisychain/dashboard/server.py +18 -5
- daisychain/spikewhale_panel.py +17 -2
- daisychain/spikewhale_task.py +4 -2
- daisychain/verified/common.py +16 -0
- daisychain/verified/mul8.py +16 -7
- daisychain/verified/ops.py +82 -20
- daisychain/verified/qat.py +126 -116
- daisychain/verified/weights/requant16.pt +1 -1
- export_luts_web.py +30 -5
- test_verified_units.py +22 -3
- web/TEST_RESULTS.md +2 -1
- web/package.json +1 -1
- web/public/requant_lut.bin +1 -1
- web/public/transformer.js +583 -582
- web/public/verified_core.js +519 -496
- web/test_nonfinite.js +25 -0
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 |
-
- **
|
| 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
|
| 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 (
|
| 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 =
|
| 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 =
|
| 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}"
|
| 181 |
-
|
|
|
|
| 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 |
-
|
|
|
|
|
|
|
|
|
|
| 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 |
-
|
| 67 |
-
|
| 68 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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 |
-
|
| 42 |
-
|
| 43 |
-
|
| 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 |
-
|
| 65 |
-
|
| 66 |
-
|
| 67 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 81 |
|
| 82 |
def dataset(self):
|
| 83 |
-
|
| 84 |
-
|
| 85 |
-
|
| 86 |
-
|
| 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 |
-
|
| 107 |
-
|
| 108 |
-
|
| 109 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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 |
-
|
| 33 |
-
|
| 34 |
-
|
| 35 |
-
|
| 36 |
-
|
| 37 |
-
|
| 38 |
-
|
| 39 |
-
|
| 40 |
-
|
| 41 |
-
|
| 42 |
-
|
| 43 |
-
|
| 44 |
-
|
| 45 |
-
|
| 46 |
-
|
| 47 |
-
|
| 48 |
-
|
| 49 |
-
|
| 50 |
-
|
| 51 |
-
|
| 52 |
-
|
| 53 |
-
|
| 54 |
-
|
| 55 |
-
|
| 56 |
-
|
| 57 |
-
|
| 58 |
-
|
| 59 |
-
|
| 60 |
-
|
| 61 |
-
|
| 62 |
-
|
| 63 |
-
|
| 64 |
-
|
| 65 |
-
|
| 66 |
-
|
| 67 |
-
|
| 68 |
-
|
| 69 |
-
|
| 70 |
-
|
| 71 |
-
|
| 72 |
-
|
| 73 |
-
|
| 74 |
-
|
| 75 |
-
|
| 76 |
-
|
| 77 |
-
|
| 78 |
-
""
|
| 79 |
-
|
| 80 |
-
|
| 81 |
-
|
| 82 |
-
|
| 83 |
-
|
| 84 |
-
|
| 85 |
-
|
| 86 |
-
|
| 87 |
-
|
| 88 |
-
|
| 89 |
-
|
| 90 |
-
|
| 91 |
-
|
| 92 |
-
|
| 93 |
-
|
| 94 |
-
|
| 95 |
-
|
| 96 |
-
|
| 97 |
-
|
| 98 |
-
|
| 99 |
-
|
| 100 |
-
|
| 101 |
-
|
| 102 |
-
|
| 103 |
-
|
| 104 |
-
|
| 105 |
-
|
| 106 |
-
|
| 107 |
-
|
| 108 |
-
|
| 109 |
-
|
| 110 |
-
|
| 111 |
-
|
| 112 |
-
|
| 113 |
-
|
| 114 |
-
|
| 115 |
-
|
| 116 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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:
|
| 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 |
-
#
|
| 36 |
-
|
| 37 |
-
|
| 38 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 39 |
print("exported mul_lut(int16 65536), requant_lut(int8 65536), relu_lut(int8 256)")
|
| 40 |
-
print("
|
|
|
|
| 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
|
| 62 |
x = np.arange(65536)
|
| 63 |
xs = np.where(x >= 32768, x - 65536, x)
|
| 64 |
-
|
| 65 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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
|
| 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:
|
| 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 |
-
|
| 396 |
-
|
| 397 |
-
|
| 398 |
-
|
| 399 |
-
const
|
| 400 |
-
|
| 401 |
-
|
| 402 |
-
|
| 403 |
-
|
| 404 |
-
|
| 405 |
-
|
| 406 |
-
|
| 407 |
-
|
| 408 |
-
const
|
| 409 |
-
const
|
| 410 |
-
const
|
| 411 |
-
|
| 412 |
-
|
| 413 |
-
|
| 414 |
-
|
| 415 |
-
|
| 416 |
-
|
| 417 |
-
|
| 418 |
-
|
| 419 |
-
|
| 420 |
-
|
| 421 |
-
//
|
| 422 |
-
|
| 423 |
-
|
| 424 |
-
bmm(cache.dlogits,
|
| 425 |
-
|
| 426 |
-
|
| 427 |
-
|
| 428 |
-
|
| 429 |
-
|
| 430 |
-
|
| 431 |
-
|
| 432 |
-
|
| 433 |
-
|
| 434 |
-
|
| 435 |
-
|
| 436 |
-
|
| 437 |
-
|
| 438 |
-
//
|
| 439 |
-
//
|
| 440 |
-
//
|
| 441 |
-
//
|
| 442 |
-
|
| 443 |
-
const
|
| 444 |
-
|
| 445 |
-
bmm(
|
| 446 |
-
|
| 447 |
-
|
| 448 |
-
|
| 449 |
-
|
| 450 |
-
|
| 451 |
-
bmm(
|
| 452 |
-
|
| 453 |
-
|
| 454 |
-
|
| 455 |
-
const
|
| 456 |
-
|
| 457 |
-
|
| 458 |
-
|
| 459 |
-
|
| 460 |
-
bmm(
|
| 461 |
-
|
| 462 |
-
|
| 463 |
-
|
| 464 |
-
//
|
| 465 |
-
|
| 466 |
-
const
|
| 467 |
-
const
|
| 468 |
-
const
|
| 469 |
-
const
|
| 470 |
-
const
|
| 471 |
-
const
|
| 472 |
-
const
|
| 473 |
-
|
| 474 |
-
bmmB(dchb,
|
| 475 |
-
|
| 476 |
-
|
| 477 |
-
//
|
| 478 |
-
|
| 479 |
-
|
| 480 |
-
|
| 481 |
-
|
| 482 |
-
|
| 483 |
-
|
| 484 |
-
for (let tj = 0; tj <= ti; tj++)
|
| 485 |
-
|
| 486 |
-
|
| 487 |
-
|
| 488 |
-
|
| 489 |
-
const
|
| 490 |
-
|
| 491 |
-
bmmB(
|
| 492 |
-
|
| 493 |
-
|
| 494 |
-
|
| 495 |
-
scatterHeadsAcc(
|
| 496 |
-
scatterHeadsAcc(
|
| 497 |
-
|
| 498 |
-
//
|
| 499 |
-
//
|
| 500 |
-
//
|
| 501 |
-
//
|
| 502 |
-
|
| 503 |
-
const
|
| 504 |
-
|
| 505 |
-
bmmB(cat(
|
| 506 |
-
|
| 507 |
-
|
| 508 |
-
|
| 509 |
-
g[gi[`b${l}.
|
| 510 |
-
g[gi[`b${l}.
|
| 511 |
-
|
| 512 |
-
//
|
| 513 |
-
//
|
| 514 |
-
//
|
| 515 |
-
//
|
| 516 |
-
|
| 517 |
-
|
| 518 |
-
|
| 519 |
-
|
| 520 |
-
|
| 521 |
-
|
| 522 |
-
|
| 523 |
-
|
| 524 |
-
|
| 525 |
-
|
| 526 |
-
|
| 527 |
-
|
| 528 |
-
|
| 529 |
-
|
| 530 |
-
|
| 531 |
-
|
| 532 |
-
|
| 533 |
-
|
| 534 |
-
|
| 535 |
-
|
| 536 |
-
|
| 537 |
-
|
| 538 |
-
const {
|
| 539 |
-
const
|
| 540 |
-
|
| 541 |
-
|
| 542 |
-
|
| 543 |
-
|
| 544 |
-
|
| 545 |
-
|
| 546 |
-
|
| 547 |
-
|
| 548 |
-
|
| 549 |
-
|
| 550 |
-
|
| 551 |
-
|
| 552 |
-
|
| 553 |
-
|
| 554 |
-
|
| 555 |
-
|
| 556 |
-
|
| 557 |
-
|
| 558 |
-
|
| 559 |
-
|
| 560 |
-
|
| 561 |
-
|
| 562 |
-
|
| 563 |
-
|
| 564 |
-
const
|
| 565 |
-
|
| 566 |
-
|
| 567 |
-
const
|
| 568 |
-
|
| 569 |
-
|
| 570 |
-
|
| 571 |
-
|
| 572 |
-
|
| 573 |
-
|
| 574 |
-
|
| 575 |
-
|
| 576 |
-
|
| 577 |
-
|
| 578 |
-
|
| 579 |
-
|
| 580 |
-
|
| 581 |
-
|
| 582 |
-
|
|
|
|
|
|
| 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 |
-
|
| 12 |
-
|
| 13 |
-
|
| 14 |
-
|
| 15 |
-
|
| 16 |
-
|
| 17 |
-
|
| 18 |
-
|
| 19 |
-
|
| 20 |
-
|
| 21 |
-
|
| 22 |
-
|
| 23 |
-
|
| 24 |
-
|
| 25 |
-
|
| 26 |
-
|
| 27 |
-
|
| 28 |
-
|
| 29 |
-
|
| 30 |
-
|
| 31 |
-
|
| 32 |
-
|
| 33 |
-
|
| 34 |
-
|
| 35 |
-
|
| 36 |
-
|
| 37 |
-
|
| 38 |
-
//
|
| 39 |
-
function
|
| 40 |
-
const
|
| 41 |
-
|
| 42 |
-
|
| 43 |
-
|
| 44 |
-
|
| 45 |
-
|
| 46 |
-
|
| 47 |
-
|
| 48 |
-
|
| 49 |
-
|
| 50 |
-
|
| 51 |
-
|
| 52 |
-
|
| 53 |
-
|
| 54 |
-
|
| 55 |
-
|
| 56 |
-
|
| 57 |
-
|
| 58 |
-
|
| 59 |
-
const
|
| 60 |
-
const
|
| 61 |
-
|
| 62 |
-
|
| 63 |
-
|
| 64 |
-
|
| 65 |
-
|
| 66 |
-
|
| 67 |
-
|
| 68 |
-
|
| 69 |
-
|
| 70 |
-
|
| 71 |
-
|
| 72 |
-
|
| 73 |
-
|
| 74 |
-
|
| 75 |
-
|
| 76 |
-
|
| 77 |
-
|
| 78 |
-
|
| 79 |
-
|
| 80 |
-
|
| 81 |
-
|
| 82 |
-
|
| 83 |
-
|
| 84 |
-
|
| 85 |
-
|
| 86 |
-
|
| 87 |
-
|
| 88 |
-
|
| 89 |
-
|
| 90 |
-
|
| 91 |
-
|
| 92 |
-
|
| 93 |
-
|
| 94 |
-
|
| 95 |
-
|
| 96 |
-
|
| 97 |
-
|
| 98 |
-
|
| 99 |
-
}
|
| 100 |
-
|
| 101 |
-
|
| 102 |
-
|
| 103 |
-
|
| 104 |
-
|
| 105 |
-
|
| 106 |
-
|
| 107 |
-
|
| 108 |
-
|
| 109 |
-
|
| 110 |
-
|
| 111 |
-
|
| 112 |
-
|
| 113 |
-
|
| 114 |
-
|
| 115 |
-
|
| 116 |
-
|
| 117 |
-
|
| 118 |
-
|
| 119 |
-
|
| 120 |
-
|
| 121 |
-
|
| 122 |
-
|
| 123 |
-
|
| 124 |
-
|
| 125 |
-
|
| 126 |
-
|
| 127 |
-
|
| 128 |
-
|
| 129 |
-
|
| 130 |
-
|
| 131 |
-
|
| 132 |
-
|
| 133 |
-
|
| 134 |
-
|
| 135 |
-
|
| 136 |
-
|
| 137 |
-
|
| 138 |
-
|
| 139 |
-
|
| 140 |
-
|
| 141 |
-
|
| 142 |
-
|
| 143 |
-
|
| 144 |
-
|
| 145 |
-
|
| 146 |
-
|
| 147 |
-
|
| 148 |
-
|
| 149 |
-
|
| 150 |
-
|
| 151 |
-
|
| 152 |
-
|
| 153 |
-
//
|
| 154 |
-
|
| 155 |
-
|
| 156 |
-
|
| 157 |
-
|
| 158 |
-
|
| 159 |
-
|
| 160 |
-
|
| 161 |
-
|
| 162 |
-
|
| 163 |
-
|
| 164 |
-
|
| 165 |
-
|
| 166 |
-
const
|
| 167 |
-
|
| 168 |
-
|
| 169 |
-
|
| 170 |
-
|
| 171 |
-
|
| 172 |
-
|
| 173 |
-
|
| 174 |
-
|
| 175 |
-
|
| 176 |
-
|
| 177 |
-
|
| 178 |
-
|
| 179 |
-
|
| 180 |
-
|
| 181 |
-
|
| 182 |
-
|
| 183 |
-
|
| 184 |
-
|
| 185 |
-
|
| 186 |
-
|
| 187 |
-
|
| 188 |
-
|
| 189 |
-
|
| 190 |
-
|
| 191 |
-
|
| 192 |
-
|
| 193 |
-
|
| 194 |
-
|
| 195 |
-
|
| 196 |
-
|
| 197 |
-
|
| 198 |
-
|
| 199 |
-
|
| 200 |
-
|
| 201 |
-
|
| 202 |
-
|
| 203 |
-
|
| 204 |
-
|
| 205 |
-
|
| 206 |
-
|
| 207 |
-
|
| 208 |
-
|
| 209 |
-
|
| 210 |
-
|
| 211 |
-
|
| 212 |
-
|
| 213 |
-
|
| 214 |
-
|
| 215 |
-
|
| 216 |
-
|
| 217 |
-
|
| 218 |
-
|
| 219 |
-
|
| 220 |
-
|
| 221 |
-
|
| 222 |
-
|
| 223 |
-
|
| 224 |
-
|
| 225 |
-
|
| 226 |
-
|
| 227 |
-
|
| 228 |
-
|
| 229 |
-
|
| 230 |
-
|
| 231 |
-
|
| 232 |
-
|
| 233 |
-
|
| 234 |
-
|
| 235 |
-
|
| 236 |
-
|
| 237 |
-
|
| 238 |
-
|
| 239 |
-
|
| 240 |
-
|
| 241 |
-
|
| 242 |
-
|
| 243 |
-
|
| 244 |
-
|
| 245 |
-
|
| 246 |
-
|
| 247 |
-
|
| 248 |
-
|
| 249 |
-
|
| 250 |
-
|
| 251 |
-
|
| 252 |
-
|
| 253 |
-
|
| 254 |
-
|
| 255 |
-
|
| 256 |
-
|
| 257 |
-
|
| 258 |
-
//
|
| 259 |
-
|
| 260 |
-
|
| 261 |
-
|
| 262 |
-
|
| 263 |
-
|
| 264 |
-
|
| 265 |
-
|
| 266 |
-
|
| 267 |
-
|
| 268 |
-
|
| 269 |
-
|
| 270 |
-
|
| 271 |
-
|
| 272 |
-
|
| 273 |
-
|
| 274 |
-
|
| 275 |
-
|
| 276 |
-
|
| 277 |
-
|
| 278 |
-
|
| 279 |
-
|
| 280 |
-
|
| 281 |
-
|
| 282 |
-
|
| 283 |
-
|
| 284 |
-
|
| 285 |
-
|
| 286 |
-
|
| 287 |
-
|
| 288 |
-
|
| 289 |
-
|
| 290 |
-
|
| 291 |
-
|
| 292 |
-
|
| 293 |
-
|
| 294 |
-
|
| 295 |
-
|
| 296 |
-
|
| 297 |
-
|
| 298 |
-
|
| 299 |
-
|
| 300 |
-
|
| 301 |
-
|
| 302 |
-
|
| 303 |
-
|
| 304 |
-
|
| 305 |
-
|
| 306 |
-
|
| 307 |
-
|
| 308 |
-
|
| 309 |
-
|
| 310 |
-
|
| 311 |
-
|
| 312 |
-
|
| 313 |
-
|
| 314 |
-
|
| 315 |
-
|
| 316 |
-
|
| 317 |
-
|
| 318 |
-
|
| 319 |
-
|
| 320 |
-
|
| 321 |
-
|
| 322 |
-
|
| 323 |
-
|
| 324 |
-
const
|
| 325 |
-
const
|
| 326 |
-
|
| 327 |
-
|
| 328 |
-
|
| 329 |
-
|
| 330 |
-
|
| 331 |
-
|
| 332 |
-
|
| 333 |
-
|
| 334 |
-
|
| 335 |
-
|
| 336 |
-
|
| 337 |
-
const
|
| 338 |
-
|
| 339 |
-
|
| 340 |
-
|
| 341 |
-
|
| 342 |
-
|
| 343 |
-
|
| 344 |
-
|
| 345 |
-
|
| 346 |
-
|
| 347 |
-
|
| 348 |
-
|
| 349 |
-
|
| 350 |
-
|
| 351 |
-
|
| 352 |
-
|
| 353 |
-
|
| 354 |
-
|
| 355 |
-
|
| 356 |
-
for (let
|
| 357 |
-
|
| 358 |
-
|
| 359 |
-
|
| 360 |
-
|
| 361 |
-
|
| 362 |
-
|
| 363 |
-
|
| 364 |
-
|
| 365 |
-
|
| 366 |
-
|
| 367 |
-
|
| 368 |
-
|
| 369 |
-
|
| 370 |
-
|
| 371 |
-
|
| 372 |
-
|
| 373 |
-
|
| 374 |
-
|
| 375 |
-
|
| 376 |
-
|
| 377 |
-
|
| 378 |
-
|
| 379 |
-
|
| 380 |
-
|
| 381 |
-
|
| 382 |
-
|
| 383 |
-
|
| 384 |
-
|
| 385 |
-
|
| 386 |
-
|
| 387 |
-
|
| 388 |
-
|
| 389 |
-
|
| 390 |
-
|
| 391 |
-
|
| 392 |
-
|
| 393 |
-
|
| 394 |
-
|
| 395 |
-
|
| 396 |
-
|
| 397 |
-
|
| 398 |
-
|
| 399 |
-
|
| 400 |
-
|
| 401 |
-
|
| 402 |
-
//
|
| 403 |
-
|
| 404 |
-
|
| 405 |
-
const
|
| 406 |
-
const
|
| 407 |
-
for (let bi = 0; bi < B; bi++)
|
| 408 |
-
|
| 409 |
-
|
| 410 |
-
|
| 411 |
-
|
| 412 |
-
|
| 413 |
-
|
| 414 |
-
|
| 415 |
-
|
| 416 |
-
|
| 417 |
-
|
| 418 |
-
|
| 419 |
-
|
| 420 |
-
|
| 421 |
-
|
| 422 |
-
|
| 423 |
-
|
| 424 |
-
|
| 425 |
-
|
| 426 |
-
|
| 427 |
-
|
| 428 |
-
|
| 429 |
-
|
| 430 |
-
|
| 431 |
-
|
| 432 |
-
|
| 433 |
-
|
| 434 |
-
|
| 435 |
-
|
| 436 |
-
|
| 437 |
-
|
| 438 |
-
|
| 439 |
-
|
| 440 |
-
|
| 441 |
-
|
| 442 |
-
|
| 443 |
-
|
| 444 |
-
|
| 445 |
-
|
| 446 |
-
const
|
| 447 |
-
const out = new Float32Array(
|
| 448 |
-
|
| 449 |
-
|
| 450 |
-
let
|
| 451 |
-
|
| 452 |
-
|
| 453 |
-
|
| 454 |
-
|
| 455 |
-
|
| 456 |
-
|
| 457 |
-
|
| 458 |
-
|
| 459 |
-
|
| 460 |
-
|
| 461 |
-
|
| 462 |
-
|
| 463 |
-
|
| 464 |
-
|
| 465 |
-
|
| 466 |
-
|
| 467 |
-
|
| 468 |
-
|
| 469 |
-
|
| 470 |
-
const
|
| 471 |
-
const
|
| 472 |
-
|
| 473 |
-
|
| 474 |
-
|
| 475 |
-
|
| 476 |
-
|
| 477 |
-
|
| 478 |
-
|
| 479 |
-
|
| 480 |
-
|
| 481 |
-
|
| 482 |
-
|
| 483 |
-
|
| 484 |
-
|
| 485 |
-
|
| 486 |
-
for (let i = 0; i <
|
| 487 |
-
|
| 488 |
-
|
| 489 |
-
|
| 490 |
-
|
| 491 |
-
|
| 492 |
-
|
| 493 |
-
|
| 494 |
-
|
| 495 |
-
|
| 496 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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.");
|