File size: 12,353 Bytes
30bafb7
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
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
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
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
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
// Forward-only transformer, split at layer boundaries, for real HF models.
//
// Every matrix product goes through the verified INT8 units — the same
// block-scaled quantize -> exact LUT/DP4A multiply -> exact int32 accumulate ->
// pinned f32 epilogue as the rest of DaisyChain. Nothing here computes in
// plain float except the elementwise parts that have no matmul in them
// (norms, softmax, RoPE, the activation), exactly as the trainer does it.
//
// The pass is cut into three callable pieces so they can run on different
// machines:
//
//   embed(ids)        head stage: token (+ learned position) embedding
//   runLayers(x)      any stage:  its own contiguous layers
//   readout(x)        head stage: final norm + lm_head -> logits
//
// Architecture differences live in the spec (see arch.js), never in branches
// scattered through the maths: `spec.norm`, `spec.gated`, `spec.rope`,
// `spec.kvHeads`. A Llama block and a GPT-2 block take the same path here.
(function (root) {
  "use strict";

  let TC, V, SH;

  // ---- norms ------------------------------------------------------------------
  // RMSNorm (Llama): x * w / sqrt(mean(x^2) + eps) — no mean subtraction.
  // LayerNorm (GPT-2): (x - mu) / sigma * w + b.
  // Both are elementwise after a row reduction, so there is no matmul to send
  // through the units, and both run in f64-accumulated JS — which IEEE requires
  // to be exactly rounded, so every device agrees.
  function rmsNorm(x, rows, C, w, eps) {
    const y = new Float32Array(rows * C);
    for (let r = 0; r < rows; r++) {
      let s = 0;
      for (let j = 0; j < C; j++) { const v = x[r * C + j]; s += v * v; }
      const inv = 1 / Math.sqrt(s / C + eps);
      for (let j = 0; j < C; j++) y[r * C + j] = x[r * C + j] * inv * w[j];
    }
    return y;
  }
  function layerNorm(x, rows, C, w, b, eps) {
    const y = new Float32Array(rows * C);
    for (let r = 0; r < rows; r++) {
      let mu = 0;
      for (let j = 0; j < C; j++) mu += x[r * C + j];
      mu /= C;
      let v = 0;
      for (let j = 0; j < C; j++) { const d = x[r * C + j] - mu; v += d * d; }
      const inv = 1 / Math.sqrt(v / C + eps);
      for (let j = 0; j < C; j++) y[r * C + j] = (x[r * C + j] - mu) * inv * w[j] + (b ? b[j] : 0);
    }
    return y;
  }
  function norm(spec, x, rows, C, w, b) {
    return spec.norm === "rms" ? rmsNorm(x, rows, C, w, spec.normEps)
                               : layerNorm(x, rows, C, w, b, spec.normEps);
  }

  // ---- activations ------------------------------------------------------------
  const silu = (v) => v / (1 + Math.exp(-v));
  // tanh-approximate GELU: what GPT-2 was trained with, so it is the correct
  // function here rather than an approximation of the erf form.
  function gelu(v) {
    return 0.5 * v * (1 + Math.tanh(0.7978845608028654 * (v + 0.044715 * v * v * v)));
  }

  // ---- RoPE -------------------------------------------------------------------
  // Rotary embeddings applied per head, in the half-split layout HF uses:
  // dims [0, hd/2) pair with [hd/2, hd). Positions are absolute, and because
  // there is no KV cache every token re-runs the whole window, so position i
  // is simply the row index.
  function ropeTables(hd, theta, maxT) {
    const half = hd >> 1;
    const cos = new Float32Array(maxT * half), sin = new Float32Array(maxT * half);
    for (let p = 0; p < maxT; p++)
      for (let i = 0; i < half; i++) {
        const f = p / Math.pow(theta, (2 * i) / hd);
        cos[p * half + i] = Math.cos(f);
        sin[p * half + i] = Math.sin(f);
      }
    return { cos, sin, half };
  }
  // x is rows x (nHeads*hd), rows = T; rotate each head in place.
  function applyRope(x, T, nHeads, hd, rope, posOffset) {
    const half = rope.half, stride = nHeads * hd;
    for (let t = 0; t < T; t++) {
      const p = (posOffset || 0) + t;
      for (let h = 0; h < nHeads; h++) {
        const o = t * stride + h * hd;
        for (let i = 0; i < half; i++) {
          const c = rope.cos[p * half + i], s = rope.sin[p * half + i];
          const a = x[o + i], b = x[o + half + i];
          x[o + i] = a * c - b * s;
          x[o + half + i] = b * c + a * s;
        }
      }
    }
  }

  // ---- verified matmul --------------------------------------------------------
  // X is m x k, W is k x n (row-major). Weights are stored k x n at load time
  // (see loadStage), so nothing is transposed per call.
  async function vmm(X, W, m, k, n, ctx, bias) {
    const out = await V.vgemmBlock(X, W, { m, k, n, batch: 1 }, ctx.L, ctx.bgemm, ctx.audit);
    if (bias) for (let i = 0; i < m; i++) for (let j = 0; j < n; j++) out[i * n + j] += bias[j];
    return out;
  }

  // ---- one layer --------------------------------------------------------------
  async function layerForward(S, x, ly) {
    const { spec, ctx } = S, C = spec.hidden, T = S.T;
    const hd = spec.headDim, nH = spec.heads, nKV = spec.kvHeads;
    const qDim = nH * hd, kvDim = nKV * hd;

    // ---- attention
    const h1 = norm(spec, x, T, C, ly.nrm1, ly.nrm1b);
    let q, k, v;
    if (spec.qkvFused) {
      // GPT-2 packs q,k,v into one (C x 3C) projection — one GEMM, then split.
      const qkv = await vmm(h1, ly.Wqkv, T, C, 3 * C, ctx, ly.bqkv);
      q = new Float32Array(T * C); k = new Float32Array(T * C); v = new Float32Array(T * C);
      for (let t = 0; t < T; t++) {
        q.set(qkv.subarray(t * 3 * C, t * 3 * C + C), t * C);
        k.set(qkv.subarray(t * 3 * C + C, t * 3 * C + 2 * C), t * C);
        v.set(qkv.subarray(t * 3 * C + 2 * C, t * 3 * C + 3 * C), t * C);
      }
    } else {
      [q, k, v] = await Promise.all([
        vmm(h1, ly.Wq, T, C, qDim, ctx, ly.bq),
        vmm(h1, ly.Wk, T, C, kvDim, ctx, ly.bk),
        vmm(h1, ly.Wv, T, C, kvDim, ctx, ly.bv),
      ]);
    }
    if (spec.rope) { applyRope(q, T, nH, hd, S.rope, 0); applyRope(k, T, nKV, hd, S.rope, 0); }

    // Grouped-query attention: several query heads share one kv head. Rather
    // than materialising repeated kv (memory this project does not have to
    // spare), the head index is mapped when reading.
    const group = nH / nKV;
    const scale = 1 / Math.sqrt(hd);
    const ctxOut = new Float32Array(T * qDim);
    // Scores and context are the two attention GEMMs. They are small per head
    // (T x hd and T x T), and the fused kernels in webgpu.js expect the
    // trainer's uniform-head layout, so at general head geometry they are run
    // through the same verified block GEMM per head — every product still goes
    // through the units, and the CPU mirror and GPU kernel agree bit-for-bit.
    for (let h = 0; h < nH; h++) {
      const kvh = Math.floor(h / group);
      const qh = new Float32Array(T * hd), kh = new Float32Array(T * hd), vh = new Float32Array(T * hd);
      for (let t = 0; t < T; t++) {
        qh.set(q.subarray(t * qDim + h * hd, t * qDim + h * hd + hd), t * hd);
        kh.set(k.subarray(t * kvDim + kvh * hd, t * kvDim + kvh * hd + hd), t * hd);
        vh.set(v.subarray(t * kvDim + kvh * hd, t * kvDim + kvh * hd + hd), t * hd);
      }
      const khT = TC.transpose(kh, T, hd);                       // hd x T
      const scores = await V.vgemmBlock(qh, khT, { m: T, k: hd, n: T, batch: 1 }, ctx.L, ctx.bgemm, ctx.audit);
      const a = new Float32Array(T * T);                          // causal softmax
      for (let i = 0; i < T; i++) {
        let mx = -Infinity;
        for (let j = 0; j <= i; j++) mx = Math.max(mx, scores[i * T + j] * scale);
        let z = 0;
        for (let j = 0; j <= i; j++) { const e = Math.exp(scores[i * T + j] * scale - mx); a[i * T + j] = e; z += e; }
        for (let j = 0; j <= i; j++) a[i * T + j] /= z;
      }
      const oh = await V.vgemmBlock(a, vh, { m: T, k: T, n: hd, batch: 1 }, ctx.L, ctx.bgemm, ctx.audit);
      for (let t = 0; t < T; t++) ctxOut.set(oh.subarray(t * hd, t * hd + hd), t * qDim + h * hd);
    }
    const attn = await vmm(ctxOut, ly.Wo, T, qDim, C, ctx, ly.bo);
    const x2 = new Float32Array(T * C);
    for (let i = 0; i < x2.length; i++) x2[i] = x[i] + attn[i];

    // ---- MLP
    const h2 = norm(spec, x2, T, C, ly.nrm2, ly.nrm2b);
    let hid;
    if (spec.gated) {
      // SwiGLU: down(silu(gate(x)) * up(x))
      const [g, u] = await Promise.all([
        vmm(h2, ly.Wgate, T, C, spec.inter, ctx, null),
        vmm(h2, ly.Wup, T, C, spec.inter, ctx, null),
      ]);
      hid = g;
      for (let i = 0; i < hid.length; i++) hid[i] = silu(hid[i]) * u[i];
    } else {
      hid = await vmm(h2, ly.Wfc, T, C, spec.inter, ctx, ly.bfc);
      for (let i = 0; i < hid.length; i++) hid[i] = gelu(hid[i]);
    }
    const down = await vmm(hid, ly.Wdown, T, spec.inter, C, ctx, ly.bdown);
    const out = new Float32Array(T * C);
    for (let i = 0; i < out.length; i++) out[i] = x2[i] + down[i];
    return out;
  }

  // ---- stage ------------------------------------------------------------------
  function makeStage(spec, st, w, ctx, T) {
    const S = { spec, st, w, ctx, T: T || Math.min(spec.maxPos, 128) };
    if (spec.rope) S.rope = ropeTables(spec.headDim, spec.ropeTheta, S.T + 1);
    return S;
  }

  function embed(S, ids) {
    const { spec } = S, C = spec.hidden, T = S.T;
    const x = new Float32Array(T * C);
    for (let i = 0; i < T; i++) {
      const id = ids[i];
      for (let j = 0; j < C; j++) x[i * C + j] = S.w.emb[id * C + j];
    }
    if (S.w.pos)                                    // GPT-2 learned positions
      for (let i = 0; i < T; i++)
        for (let j = 0; j < C; j++) x[i * C + j] += S.w.pos[i * C + j];
    return x;
  }

  async function runLayers(S, x) {
    for (const ly of S.w.layers) x = await layerForward(S, x, ly);
    return x;
  }

  // Only the last position's logits are needed to pick the next token, and at
  // a 150k-token vocabulary that turns the largest GEMM in the model into a
  // single row — the biggest single saving in the ring.
  async function readout(S, x) {
    const { spec, ctx } = S, C = spec.hidden, T = S.T;
    const y = norm(spec, x, T, C, S.w.nrmF, S.w.nrmFb);
    const last = y.subarray((T - 1) * C, T * C);
    return vmm(last, S.w.lmHead, 1, C, spec.vocab, ctx, null);
  }

  // ---- sampling ---------------------------------------------------------------
  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; }; }
  function pickToken(logits, opts) {
    const o = opts || {}, temp = o.temperature ?? 0, vocab = logits.length;
    if (!temp) {
      let best = 0, bv = -Infinity;
      for (let j = 0; j < vocab; j++) if (logits[j] > bv) { bv = logits[j]; best = j; }
      return best;
    }
    const k = Math.min(o.topK || 40, vocab);
    const idx = Array.from({ length: vocab }, (_, i) => i).sort((a, b) => logits[b] - logits[a]).slice(0, k);
    let mx = -Infinity;
    for (const i of idx) mx = Math.max(mx, logits[i] / temp);
    let z = 0;
    const p = idx.map(i => { const e = Math.exp(logits[i] / temp - mx); z += e; return e; });
    let r = (o.rng || Math.random)() * z;
    for (let i = 0; i < idx.length; i++) { r -= p[i]; if (r <= 0) return idx[i]; }
    return idx[idx.length - 1];
  }

  // ---- single-device reference ------------------------------------------------
  async function generateLocal(stages, ids, nTokens, opts) {
    const head = stages[0], T = head.T;
    const out = [...ids];
    for (let n = 0; n < nTokens; n++) {
      const win = new Int32Array(T);
      const tail = out.slice(-T);
      for (let i = 0; i < tail.length; i++) win[T - tail.length + i] = tail[i];
      let x = embed(head, win);
      for (const S of stages) x = await runLayers(S, x);
      out.push(pickToken(await readout(head, x), opts));
    }
    return out;
  }

  const api = { makeStage, embed, runLayers, readout, pickToken, generateLocal,
                rmsNorm, layerNorm, norm, ropeTables, applyRope, silu, gelu, mulberry32 };
  if (typeof module !== "undefined" && module.exports) {
    TC = require("./traincore.js"); V = require("./verified_core.js"); SH = require("./shard.js");
    module.exports = api;
  } else { TC = root.TrainCore; V = root.Verified; SH = root.Shard; root.Infer = api; }
})(typeof self !== "undefined" ? self : this);