Spaces:
Running
Running
spec-decode for BitNet: subNorm+f32-KV batched verify, ?spec + ?bench=spec
Browse files- qvac-gpu.js +14 -9
qvac-gpu.js
CHANGED
|
@@ -1826,19 +1826,20 @@ export async function createQvacGPU(manifest, fetchTensor, cap = 64, eos = 2, st
|
|
| 1826 |
let _drafter = null; // pluggable learned drafter: (seq, max) => ids (sync or async). null = n-gram baseline.
|
| 1827 |
function specInit() {
|
| 1828 |
if (SP) return SP;
|
| 1829 |
-
if (
|
| 1830 |
const sb8 = (n) => dev.createBuffer({ size: n * 4, usage: U.STORAGE | U.COPY_DST | U.COPY_SRC });
|
| 1831 |
SP = {
|
| 1832 |
x: sb8(KX * d), normed: sb8(KX * d), q: sb8(KX * q_dim), k: sb8(KX * kv_dim), v: sb8(KX * kv_dim),
|
| 1833 |
-
attn: sb8(KX * q_dim), h: sb8(KX * d), normed2: sb8(KX * d), gate: sb8(KX * ff), up: sb8(KX * ff),
|
| 1834 |
-
hid: sb8(KX * ff), cur: sb8(KX * d), logits: sb8(KX * vocab), amax: sb8(KX), amaxStg: dev.createBuffer({ size: KX * 4, usage: U.MAP_READ | U.COPY_DST }),
|
| 1835 |
tmp: sbuf(65536 * 2),
|
| 1836 |
P_rmsk: pipe(RMSK, "rmsk"), P_ropek: pipe(ROPEK(ropeLit), "ropek"), P_kvqk: pipe(KVQK(kv_dim), "kvqk"),
|
| 1837 |
P_attnqk: pipe(ATTNQK(cap, kv_dim), "attnqk"), P_attnk: pipe(ATTNK(cap), "attnk"),
|
| 1838 |
P_t2k: pipe(mmT2KK(false, KX), "t2k"), P_t2ka: pipe(mmT2KK(true, KX), "t2ka"),
|
| 1839 |
P_t2k2: pipe(mmT2KK(false, 2), "t2k2"), P_t2ka2: pipe(mmT2KK(true, 2), "t2ka2"), P_q3k2: pipe(mmQ3KK(2), "q3k2"),
|
| 1840 |
P_q3k: pipe(mmQ3KK(KX), "q3k"),
|
| 1841 |
-
uD: ubuf(new Uint32Array([d, 0, 0, 0])),
|
|
|
|
| 1842 |
uAT: ubuf(new Uint32Array([4])), uFFK: ubuf(new Uint32Array([4])),
|
| 1843 |
uPen: Array.from({ length: KX }, () => ubuf(new Uint32Array(4))), uRow: Array.from({ length: KX }, (_, i) => ubuf(new Uint32Array([i, 0, 0, 0]))), uV: ubuf(new Uint32Array([vocab, 0, 0, 0])),
|
| 1844 |
cold: 0, stats: { windows: 0, drafted: 0, accepted: 0 },
|
|
@@ -1912,19 +1913,23 @@ export async function createQvacGPU(manifest, fetchTensor, cap = 64, eos = 2, st
|
|
| 1912 |
specMM(enc, W_("wq"), SP.normed, SP.q, m); specMM(enc, W_("wk"), SP.normed, SP.k, m); specMM(enc, W_("wv"), SP.normed, SP.v, m);
|
| 1913 |
pass(enc, SP.P_ropek, [SP.q, SP.uRQ], [Math.ceil(n_heads * (hd / 2) / 64), m]);
|
| 1914 |
pass(enc, SP.P_ropek, [SP.k, SP.uRK], [Math.ceil(n_kv_heads * (hd / 2) / 64), m]);
|
| 1915 |
-
if (l === 0) {
|
| 1916 |
-
for (let i = 0; i < m; i++) { enc.copyBufferToBuffer(SP.k, i * kv_dim * 4, kcache[
|
| 1917 |
-
pass(enc, SP.P_attnk, [SP.q, kcache[
|
| 1918 |
} else {
|
| 1919 |
pass(enc, SP.P_kvqk, [SP.k, kcache[l], SP.uAT], [1, m]);
|
| 1920 |
pass(enc, SP.P_kvqk, [SP.v, vcache[l], SP.uAT], [1, m]);
|
| 1921 |
pass(enc, SP.P_attnqk, [SP.q, kcache[l], vcache[l], SP.attn, SP.uAT], [n_heads, m]);
|
| 1922 |
}
|
| 1923 |
-
|
|
|
|
|
|
|
| 1924 |
pass(enc, SP.P_rmsk, [SP.h, Nrm[`l${l}.ffn_norm`], SP.normed2, SP.uD], [1, m]);
|
| 1925 |
specMM(enc, W_("w_gate"), SP.normed2, SP.gate, m); specMM(enc, W_("w_up"), SP.normed2, SP.up, m);
|
| 1926 |
pass(enc, P_sm, [SP.gate, SP.up, SP.hid, SP.uFFK], Math.ceil(m * ff / 64));
|
| 1927 |
-
|
|
|
|
|
|
|
| 1928 |
cur = SP.cur; if (l < n_layers - 1) { const t_ = SP.x; SP.x = SP.cur; SP.cur = t_; }
|
| 1929 |
}
|
| 1930 |
pass(enc, SP.P_rmsk, [cur, Nrm["final_norm"], SP.normed, SP.uD], [1, m]);
|
|
|
|
| 1826 |
let _drafter = null; // pluggable learned drafter: (seq, max) => ids (sync or async). null = n-gram baseline.
|
| 1827 |
function specInit() {
|
| 1828 |
if (SP) return SP;
|
| 1829 |
+
if (bitlinear || moe || stream || frameGran || !hasT2) throw new Error("specDecode: ternary (t2) resident non-MoE models only (for now)");
|
| 1830 |
const sb8 = (n) => dev.createBuffer({ size: n * 4, usage: U.STORAGE | U.COPY_DST | U.COPY_SRC });
|
| 1831 |
SP = {
|
| 1832 |
x: sb8(KX * d), normed: sb8(KX * d), q: sb8(KX * q_dim), k: sb8(KX * kv_dim), v: sb8(KX * kv_dim),
|
| 1833 |
+
attn: sb8(KX * q_dim), attn2: sb8(KX * q_dim), h: sb8(KX * d), normed2: sb8(KX * d), gate: sb8(KX * ff), up: sb8(KX * ff),
|
| 1834 |
+
hid: sb8(KX * ff), hid2: sb8(KX * ff), cur: sb8(KX * d), logits: sb8(KX * vocab), amax: sb8(KX), amaxStg: dev.createBuffer({ size: KX * 4, usage: U.MAP_READ | U.COPY_DST }),
|
| 1835 |
tmp: sbuf(65536 * 2),
|
| 1836 |
P_rmsk: pipe(RMSK, "rmsk"), P_ropek: pipe(ROPEK(ropeLit), "ropek"), P_kvqk: pipe(KVQK(kv_dim), "kvqk"),
|
| 1837 |
P_attnqk: pipe(ATTNQK(cap, kv_dim), "attnqk"), P_attnk: pipe(ATTNK(cap), "attnk"),
|
| 1838 |
P_t2k: pipe(mmT2KK(false, KX), "t2k"), P_t2ka: pipe(mmT2KK(true, KX), "t2ka"),
|
| 1839 |
P_t2k2: pipe(mmT2KK(false, 2), "t2k2"), P_t2ka2: pipe(mmT2KK(true, 2), "t2ka2"), P_q3k2: pipe(mmQ3KK(2), "q3k2"),
|
| 1840 |
P_q3k: pipe(mmQ3KK(KX), "q3k"),
|
| 1841 |
+
uD: ubuf(new Uint32Array([d, 0, 0, 0])), uQd: ubuf(new Uint32Array([q_dim, 0, 0, 0])), uFFd: ubuf(new Uint32Array([ff, 0, 0, 0])),
|
| 1842 |
+
uRQ: ubuf(new Uint32Array([4])), uRK: ubuf(new Uint32Array([4])),
|
| 1843 |
uAT: ubuf(new Uint32Array([4])), uFFK: ubuf(new Uint32Array([4])),
|
| 1844 |
uPen: Array.from({ length: KX }, () => ubuf(new Uint32Array(4))), uRow: Array.from({ length: KX }, (_, i) => ubuf(new Uint32Array([i, 0, 0, 0]))), uV: ubuf(new Uint32Array([vocab, 0, 0, 0])),
|
| 1845 |
cold: 0, stats: { windows: 0, drafted: 0, accepted: 0 },
|
|
|
|
| 1913 |
specMM(enc, W_("wq"), SP.normed, SP.q, m); specMM(enc, W_("wk"), SP.normed, SP.k, m); specMM(enc, W_("wv"), SP.normed, SP.v, m);
|
| 1914 |
pass(enc, SP.P_ropek, [SP.q, SP.uRQ], [Math.ceil(n_heads * (hd / 2) / 64), m]);
|
| 1915 |
pass(enc, SP.P_ropek, [SP.k, SP.uRK], [Math.ceil(n_kv_heads * (hd / 2) / 64), m]);
|
| 1916 |
+
if (l === 0 || !kv4) { // f32 KV cache (all layers when the model isn't kv4, e.g. BitNet)
|
| 1917 |
+
for (let i = 0; i < m; i++) { enc.copyBufferToBuffer(SP.k, i * kv_dim * 4, kcache[l], (base + i) * kv_dim * 4, kv_dim * 4); enc.copyBufferToBuffer(SP.v, i * kv_dim * 4, vcache[l], (base + i) * kv_dim * 4, kv_dim * 4); }
|
| 1918 |
+
pass(enc, SP.P_attnk, [SP.q, kcache[l], vcache[l], SP.attn, SP.uAT], [n_heads, m]);
|
| 1919 |
} else {
|
| 1920 |
pass(enc, SP.P_kvqk, [SP.k, kcache[l], SP.uAT], [1, m]);
|
| 1921 |
pass(enc, SP.P_kvqk, [SP.v, vcache[l], SP.uAT], [1, m]);
|
| 1922 |
pass(enc, SP.P_attnqk, [SP.q, kcache[l], vcache[l], SP.attn, SP.uAT], [n_heads, m]);
|
| 1923 |
}
|
| 1924 |
+
let attnO = SP.attn;
|
| 1925 |
+
if (subNorm) { pass(enc, SP.P_rmsk, [SP.attn, Nrm[`l${l}.attn_sub_norm`], SP.attn2, SP.uQd], [1, m]); attnO = SP.attn2; } // BitNet: RMSNorm before wo
|
| 1926 |
+
specMM(enc, W_("wo"), attnO, SP.h, m, cur);
|
| 1927 |
pass(enc, SP.P_rmsk, [SP.h, Nrm[`l${l}.ffn_norm`], SP.normed2, SP.uD], [1, m]);
|
| 1928 |
specMM(enc, W_("w_gate"), SP.normed2, SP.gate, m); specMM(enc, W_("w_up"), SP.normed2, SP.up, m);
|
| 1929 |
pass(enc, P_sm, [SP.gate, SP.up, SP.hid, SP.uFFK], Math.ceil(m * ff / 64));
|
| 1930 |
+
let hidO = SP.hid;
|
| 1931 |
+
if (subNorm) { pass(enc, SP.P_rmsk, [SP.hid, Nrm[`l${l}.ffn_sub_norm`], SP.hid2, SP.uFFd], [1, m]); hidO = SP.hid2; } // BitNet: RMSNorm before w_down
|
| 1932 |
+
specMM(enc, W_("w_down"), hidO, SP.cur, m, SP.h);
|
| 1933 |
cur = SP.cur; if (l < n_layers - 1) { const t_ = SP.x; SP.x = SP.cur; SP.cur = t_; }
|
| 1934 |
}
|
| 1935 |
pass(enc, SP.P_rmsk, [cur, Nrm["final_norm"], SP.normed, SP.uD], [1, m]);
|