Quazim0t0's picture
DaisyChain-Infer: project files + guide
30bafb7 verified
Raw
History Blame Contribute Delete
12.4 kB
// 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);