Humuhumu33 commited on
Commit
57800b6
·
verified ·
1 Parent(s): bbf8a0a

spec-decode for BitNet: subNorm+f32-KV batched verify, ?spec + ?bench=spec

Browse files
Files changed (1) hide show
  1. 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 (subNorm || bitlinear || fusedT2 || !kv4 || moe || stream) throw new Error("specDecode: unfused+kv4 path 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), 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])), uRQ: ubuf(new Uint32Array([4])), uRK: ubuf(new Uint32Array([4])),
 
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) { // layer 0: f32 cache
1916
- for (let i = 0; i < m; i++) { enc.copyBufferToBuffer(SP.k, i * kv_dim * 4, kcache[0], (base + i) * kv_dim * 4, kv_dim * 4); enc.copyBufferToBuffer(SP.v, i * kv_dim * 4, vcache[0], (base + i) * kv_dim * 4, kv_dim * 4); }
1917
- pass(enc, SP.P_attnk, [SP.q, kcache[0], vcache[0], SP.attn, SP.uAT], [n_heads, m]);
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
- specMM(enc, W_("wo"), SP.attn, SP.h, m, cur);
 
 
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
- specMM(enc, W_("w_down"), SP.hid, SP.cur, m, SP.h);
 
 
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]);