| <!DOCTYPE html> |
| <html lang="en"> |
| <head> |
| <meta charset="utf-8"/> |
| <meta name="viewport" content="width=device-width, initial-scale=1"/> |
| <title>Magenta RT · Live (custom WGSL temporal + custom WGSL depth)</title> |
| <style> |
| :root{ --bg:#fffdf3; --ink:#1c1917; --mut:#78716c; --amber:#f59e0b; --amber2:#fbbf24; --lime:#84cc16; --line:#e7e1c8; --card:#fffef9; } |
| *{ box-sizing:border-box; } body{ margin:0; font-family:"Segoe UI",-apple-system,sans-serif; color:var(--ink); |
| background:radial-gradient(900px 480px at 80% -5%, #fef3c7, var(--bg) 55%); min-height:100vh; } |
| .wrap{ max-width:680px; margin:0 auto; padding:30px 20px 60px; } |
| h1{ font-size:23px; margin:0 0 2px; font-weight:800; letter-spacing:-.4px; } |
| h1 .em{ background:linear-gradient(90deg,var(--amber),var(--lime)); -webkit-background-clip:text; background-clip:text; color:transparent; } |
| .sub{ color:var(--mut); font-size:13px; margin-bottom:22px; } |
| .card{ background:var(--card); border:1px solid var(--line); border-radius:16px; padding:18px; margin-bottom:14px; } |
| label{ font-size:12px; color:var(--mut); text-transform:uppercase; letter-spacing:.5px; display:block; margin-bottom:8px; } |
| .row{ display:flex; gap:18px; flex-wrap:wrap; align-items:flex-end; } |
| select,button.go{ font-family:inherit; font-weight:700; } |
| select{ padding:10px 12px; border-radius:10px; border:1.5px solid var(--line); background:#fff; font-size:15px; color:var(--ink); } |
| button.go{ padding:12px 26px; border:none; border-radius:999px; font-size:15px; color:#3a2c05; |
| background:linear-gradient(135deg,var(--amber2),var(--amber)); box-shadow:0 4px 14px #f59e0b55; cursor:pointer; } |
| button.go:disabled{ opacity:.5; cursor:default; } |
| .stat{ font-size:13px; color:var(--mut); margin-top:10px; min-height:18px; } |
| .brk{ font-size:12px; color:var(--mut); margin-top:4px; min-height:16px; font-variant-numeric:tabular-nums; } |
| canvas{ width:100%; height:70px; display:block; border-radius:10px; background:#fffdf3; margin-top:10px; border:1px solid var(--line); } |
| #prog{ display:none; height:10px; background:#fef9e7; border:1px solid var(--line); border-radius:999px; overflow:hidden; margin-top:10px; } |
| #bar{ height:100%; width:0%; border-radius:999px; background:linear-gradient(90deg,var(--amber),var(--lime)); transition:width .15s; } |
| code{ background:#fef9e7; padding:1px 6px; border-radius:6px; font-size:12px; } |
| .badge{ font-size:11px; padding:2px 8px; border-radius:999px; background:#f7fee7; color:#3f6212; border:1px solid #d9f99d; } |
| </style> |
| </head> |
| <body> |
| <div class="wrap"> |
| <h1>🍋 Magenta RT · <span class="em">full custom WGSL on your GPU</span></h1> |
| <div class="sub">ORT encoder → <b>custom WGSL temporal</b> → <b>custom WGSL depth</b> (all GPU-resident, shared ORT device) → ORT streaming decoder. Both transformers run as hand-written WebGPU kernels.</div> |
|
|
| <div class="card"> |
| <label>Style</label> |
| <div class="row"> |
| <select id="prompt"></select> |
| <button class="go" id="load">Load model</button> |
| <button class="go" id="play" disabled>▶ Play</button> |
| <button class="go" id="stop" disabled style="background:#2a2e45;color:#fff">■ Stop</button> |
| <span class="badge" id="rt"></span> |
| </div> |
| <div class="stat" id="stat">Press <b>Load model</b> (downloads encoder + decoder + 390 MB WGSL temporal + 148 MB WGSL depth weights, cached after first run).</div> |
| <div class="brk" id="brk"></div> |
| <div id="prog"><div id="bar"></div></div> |
| <canvas id="viz" width="640" height="70"></canvas> |
| </div> |
| </div> |
|
|
| <script type="module"> |
| import * as ort from "https://cdn.jsdelivr.net/npm/onnxruntime-web@1.23.0/dist/ort.webgpu.min.mjs"; |
| ort.env.wasm.numThreads = navigator.hardwareConcurrency ? Math.min(4, navigator.hardwareConcurrency) : 2; |
| ort.env.webgpu.powerPreference = "high-performance"; |
| |
| const $ = id => document.getElementById(id); |
| |
| const sz = "small", dev = "webgpu"; |
| const REPO = `https://huggingface.co/magenta-torch/magenta-rt-onnx-${sz}/resolve/main`; |
| const COMMUNITY = "https://huggingface.co/magenta-community/magenta-rt-onnx-small/resolve/main"; |
| const DEC_URL = COMMUNITY + "/spectrostream_stream.onnx"; |
| const DEC_STATE_URL = COMMUNITY + "/stream_state.json"; |
| const DEPTH_BASE = COMMUNITY + "/depth_wgsl/"; |
| const TEMP_BASE = COMMUNITY + "/temporal_wgsl/"; |
| let cfg=null, S={}, EMB=null, QUANT=null, presets={}; |
| let DSJ=null, DSTATE=null, HIFLUSH=undefined; |
| let DEPTH=null; |
| let TEMPORAL=null; |
| |
| function setBar(f){ $("bar").style.width=Math.min(100,Math.max(0,f*100)).toFixed(1)+"%"; } |
| |
| |
| function f16(u){ const s=(u&0x8000)>>15,e=(u&0x7c00)>>10,f=u&0x03ff; |
| if(e===0) return (s?-1:1)*Math.pow(2,-14)*(f/1024); |
| if(e===31) return f?NaN:(s?-1:1)*Infinity; |
| return (s?-1:1)*Math.pow(2,e-15)*(1+f/1024); } |
| function f16decode(b){ const u=new Uint16Array(b.buffer,b.byteOffset,b.byteLength>>1), o=new Float32Array(u.length); for(let i=0;i<u.length;i++)o[i]=f16(u[i]); return o; } |
| |
| |
| async function fetchBuf(url, onFrac){ |
| let cache=null; |
| try{ cache=await caches.open("mrt-fullcustom-v1"); const hit=await cache.match(url); |
| if(hit){ onFrac&&onFrac(1); return new Uint8Array(await hit.arrayBuffer()); } }catch(e){} |
| const r=await fetch(url); if(!r.ok) throw new Error(r.status+" "+url.split("/").pop()); |
| const total=+r.headers.get("content-length")||0, reader=r.body.getReader(), chunks=[]; let recv=0; |
| for(;;){ const {done,value}=await reader.read(); if(done)break; chunks.push(value); recv+=value.length; if(total&&onFrac)onFrac(recv/total); } |
| const buf=new Uint8Array(recv); let o=0; for(const c of chunks){buf.set(c,o);o+=c.length;} |
| if(cache){ try{ await cache.put(url, new Response(buf)); }catch(e){} } |
| return buf; |
| } |
| async function fetchJSON(url){ const r=await fetch(url); if(!r.ok) throw new Error(r.status+" "+url); return r.json(); } |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| const T_D_MODEL=1024, T_D_SRC=256, T_D_FF=4096, T_NUM_LAYERS=12, T_NUM_HEADS=8, |
| T_UPH=128, T_HD=T_NUM_HEADS*T_UPH, T_KV_CAP=42, T_EPS=1e-6, T_R_SOFTPLUS_0=1.442695041; |
| |
| const T_RED=32, T_ROWS=8, T_T_WG=T_ROWS*T_RED; |
| const tNWGT=N=>Math.ceil(N/T_ROWS); |
| const T_TILE_ROWS=64, T_LANES=4, T_WG2=T_TILE_ROWS*T_LANES; |
| const tNWG=N=>Math.ceil(N/T_TILE_ROWS); |
| |
| |
| function f32ToF16(val){ |
| const f=new Float32Array(1), i=new Int32Array(f.buffer); f[0]=val; const x=i[0]; |
| const sign=(x>>16)&0x8000; let mant=x&0x007fffff; let exp=(x>>23)&0xff; |
| if(exp===0xff) return sign|0x7c00|(mant?0x0200:0); |
| exp=exp-127+15; |
| if(exp>=0x1f) return sign|0x7c00; |
| if(exp<=0){ if(exp<-10) return sign; mant=mant|0x00800000; const shift=14-exp; |
| let half=mant>>shift; const rem=mant&((1<<shift)-1); const halfway=1<<(shift-1); |
| if(rem>halfway||(rem===halfway&&(half&1))) half+=1; return sign|half; } |
| let half=(exp<<10)|(mant>>13); const rem=mant&0x1fff; |
| if(rem>0x1000||(rem===0x1000&&(half&1))) half+=1; return sign|half; |
| } |
| |
| |
| |
| async function buildTemporal(device, layout, binBytes){ |
| let USE_SUBGROUPS=false; |
| const adapter = device.__mrtAdapter || null; |
| |
| const D_MODEL=T_D_MODEL, D_SRC=T_D_SRC, D_FF=T_D_FF, NUM_LAYERS=T_NUM_LAYERS, |
| NUM_HEADS=T_NUM_HEADS, UPH=T_UPH, HD=T_HD, KV_CAP=T_KV_CAP, EPS=T_EPS, |
| R_SOFTPLUS_0=T_R_SOFTPLUS_0, RED=T_RED, ROWS=T_ROWS, T_WG=T_T_WG, |
| TILE_ROWS=T_TILE_ROWS, LANES=T_LANES, WG2=T_WG2; |
| |
| const EMBED_SCALE=layout.embed_scale, TOTAL_FLOATS=layout.total_floats; |
| const off=n=>layout.tensors[n].offset; |
| const L=(li,blk,t)=>off(`L${li}.${blk}.${t}`); |
| |
| |
| const TOTAL_PACK=layout.padded_floats; |
| const STORE_WORDS=TOTAL_PACK/2; |
| const W_SPLIT_WORDS=Math.ceil(STORE_WORDS/2); |
| |
| |
| const u16=new Uint16Array(binBytes.buffer,binBytes.byteOffset,binBytes.byteLength/2); |
| const binF32=new Float32Array(TOTAL_FLOATS); |
| for(let i=0;i<TOTAL_FLOATS;i++) binF32[i]=f16(u16[i]); |
| |
| |
| const CONSTS=` |
| const D_MODEL : u32 = ${D_MODEL}u; |
| const D_SRC : u32 = ${D_SRC}u; |
| const D_FF : u32 = ${D_FF}u; |
| const NUM_HEADS : u32 = ${NUM_HEADS}u; |
| const UPH : u32 = ${UPH}u; |
| const HD : u32 = ${HD}u; |
| const EPS : f32 = ${EPS}; |
| const R_SOFTPLUS_0 : f32 = ${R_SOFTPLUS_0}; |
| const RED : u32 = ${RED}u; |
| const ROWS : u32 = ${ROWS}u; |
| const TILE_ROWS : u32 = ${TILE_ROWS}u; |
| const LANES : u32 = ${LANES}u; |
| `; |
| const WBUFS=` |
| @group(0) @binding(0) var<storage, read> W0 : array<u32>; |
| @group(0) @binding(6) var<storage, read> W1 : array<u32>; |
| const W_SPLIT_WORDS : u32 = ${W_SPLIT_WORDS}u; |
| fn wword(wi : u32) -> u32 { |
| if (wi < W_SPLIT_WORDS) { return W0[wi]; } |
| return W1[wi - W_SPLIT_WORDS]; |
| } |
| fn wv(i : u32) -> f32 { |
| let wi = i >> 1u; |
| let packed = wword(wi); |
| let two = unpack2x16float(packed); |
| if ((i & 1u) == 0u) { return two.x; } |
| return two.y; |
| } |
| `; |
| const COMMON=` |
| fn softplus(x: f32) -> f32 { |
| return log(1.0 + exp(-abs(x))) + max(x, 0.0); |
| } |
| fn gelu_tanh(x: f32) -> f32 { |
| let c : f32 = 0.7978845608028654; // sqrt(2/pi) |
| return 0.5 * x * (1.0 + tanh(c * (x + 0.044715 * x * x * x))); |
| } |
| `; |
| const HEADER=CONSTS+WBUFS+COMMON; |
| const WTDECL=` |
| @group(0) @binding(7) var<storage, read> WT : array<u32>; |
| fn wtv(i : u32) -> f32 { |
| let two = unpack2x16float(WT[i >> 1u]); |
| if ((i & 1u) == 0u) { return two.x; } |
| return two.y; |
| } |
| `; |
| |
| |
| function gridMatmulT(opts){ |
| const K=opts.K, sg=USE_SUBGROUPS; |
| const inDecl=`@group(0) @binding(${opts.inBinding}) var<storage, read> vin : array<f32>;`; |
| const outDecl=`@group(0) @binding(${opts.outBinding}) var<storage, read_write> vout : array<f32>;`; |
| const uniformDecl=`struct U { wOff:u32, colBase:u32, normOff:u32, biasOff:u32, }; |
| @group(0) @binding(5) var<uniform> u : U;`; |
| const normReduce=opts.rmsnorm?` |
| var lss : f32 = 0.0; |
| for (var i : u32 = tid; i < ${K}u; i = i + ${T_WG}u) { |
| let xv = vin[i]; lss = lss + xv * xv; |
| } |
| red[tid] = lss; |
| workgroupBarrier(); |
| for (var s : u32 = ${T_WG/2}u; s > 0u; s = s >> 1u) { |
| if (tid < s) { red[tid] = red[tid] + red[tid + s]; } |
| workgroupBarrier(); |
| } |
| let rms = 1.0 / sqrt(red[0] / f32(${K}u) + EPS); |
| workgroupBarrier(); |
| for (var i : u32 = tid; i < ${K}u; i = i + ${T_WG}u) { |
| inS[i] = (vin[i] * rms) * wv(u.normOff + i); |
| } |
| workgroupBarrier();`:` |
| for (var i : u32 = tid; i < ${K}u; i = i + ${T_WG}u) { inS[i] = vin[i]; } |
| workgroupBarrier();`; |
| let epi=""; |
| if(opts.epilogue==="geluBias"){ epi=`acc = gelu_tanh(acc + wv(u.biasOff + col));`; } |
| else { epi=opts.hasBias?`acc = acc + wv(u.biasOff + col);`:``; } |
| const reduceWrite=sg?` |
| let total = subgroupAdd(acc); |
| if (lane == 0u) { |
| var accf : f32 = total; |
| ${epi.replace(/\bacc\b/g,"accf")} |
| vout[outIdx] = accf; |
| }`:` |
| part[tid] = acc; |
| workgroupBarrier(); |
| for (var s : u32 = ${RED/2}u; s > 0u; s = s >> 1u) { |
| if (lane < s) { part[tid] = part[tid] + part[tid + s]; } |
| workgroupBarrier(); |
| } |
| if (lane == 0u) { |
| var accf : f32 = part[row * RED]; |
| ${epi.replace(/\bacc\b/g,"accf")} |
| vout[outIdx] = accf; |
| }`; |
| const enableSg=sg?"enable subgroups;\n":""; |
| return enableSg+HEADER+WTDECL+` |
| ${inDecl} |
| ${outDecl} |
| ${uniformDecl} |
| |
| var<workgroup> inS : array<f32, ${K}>; |
| ${opts.rmsnorm?`var<workgroup> red : array<f32, ${T_WG}>;`:``} |
| ${sg?"":`var<workgroup> part : array<f32, ${T_WG}>;`} |
| |
| @compute @workgroup_size(${T_WG}) |
| fn main(@builtin(local_invocation_id) lid : vec3<u32>, |
| @builtin(workgroup_id) wid : vec3<u32>) { |
| let tid = lid.x; |
| let row = tid / RED; |
| let lane = tid % RED; |
| ${normReduce} |
| let outIdx = wid.x * ROWS + row; |
| let col = u.colBase + outIdx; |
| let wbase = u.wOff + col * ${K}u; |
| var acc : f32 = 0.0; |
| for (var k : u32 = lane; k < ${K}u; k = k + RED) { |
| acc = acc + inS[k] * wtv(wbase + k); |
| } |
| ${reduceWrite} |
| } |
| `; |
| } |
| |
| |
| function buildKV(opts){ |
| const K=opts.K, sg=USE_SUBGROUPS, enableSg=sg?"enable subgroups;\n":""; |
| const qPart=opts.writeQ?`var qa : f32 = 0.0;`:``; |
| const qAccum=opts.writeQ?`qa = qa + hv * wtv(qb + k);`:``; |
| const reduceWrite=sg?` |
| let sk = subgroupAdd(ka); |
| let sv = subgroupAdd(va); |
| ${opts.writeQ?"let sq = subgroupAdd(qa);":""} |
| if (lane == 0u) { |
| let cbase = u.slot * HD; |
| kcache[cbase + oi] = sk; |
| vcache[cbase + oi] = sv; |
| ${opts.writeQ?"qbuf[oi] = sq;":""} |
| }`:` |
| pk[tid] = ka; pv[tid] = va; |
| ${opts.writeQ?"pq[tid] = qa;":""} |
| workgroupBarrier(); |
| for (var s : u32 = ${RED/2}u; s > 0u; s = s >> 1u) { |
| if (lane < s) { |
| pk[tid] = pk[tid] + pk[tid + s]; |
| pv[tid] = pv[tid] + pv[tid + s]; |
| ${opts.writeQ?"pq[tid] = pq[tid] + pq[tid + s];":""} |
| } |
| workgroupBarrier(); |
| } |
| if (lane == 0u) { |
| let cbase = u.slot * HD; |
| kcache[cbase + oi] = pk[row * RED]; |
| vcache[cbase + oi] = pv[row * RED]; |
| ${opts.writeQ?"qbuf[oi] = pq[row * RED];":""} |
| }`; |
| const normReduce=opts.rmsnorm?` |
| var lss : f32 = 0.0; |
| for (var i : u32 = tid; i < ${K}u; i = i + ${T_WG}u) { |
| let xv = vin[i]; lss = lss + xv * xv; |
| } |
| red[tid] = lss; |
| workgroupBarrier(); |
| for (var s : u32 = ${T_WG/2}u; s > 0u; s = s >> 1u) { |
| if (tid < s) { red[tid] = red[tid] + red[tid + s]; } |
| workgroupBarrier(); |
| } |
| let rms = 1.0 / sqrt(red[0] / f32(${K}u) + EPS); |
| workgroupBarrier(); |
| for (var i : u32 = tid; i < ${K}u; i = i + ${T_WG}u) { |
| inS[i] = (vin[i] * rms) * wv(u.normOff + i); |
| } |
| workgroupBarrier();`:` |
| for (var i : u32 = tid; i < ${K}u; i = i + ${T_WG}u) { inS[i] = vin[i]; } |
| workgroupBarrier();`; |
| return enableSg+HEADER+WTDECL+` |
| @group(0) @binding(1) var<storage, read> vin : array<f32>; |
| ${opts.writeQ?`@group(0) @binding(2) var<storage, read_write> qbuf : array<f32>;`:``} |
| @group(0) @binding(3) var<storage, read_write> kcache : array<f32>; |
| @group(0) @binding(4) var<storage, read_write> vcache : array<f32>; |
| struct U { normOff:u32, Wq:u32, Wk:u32, Wv:u32, slot:u32, }; |
| @group(0) @binding(5) var<uniform> u : U; |
| |
| var<workgroup> inS : array<f32, ${K}>; |
| var<workgroup> red : array<f32, ${T_WG}>; |
| ${sg?"":`var<workgroup> pk : array<f32, ${T_WG}>; |
| var<workgroup> pv : array<f32, ${T_WG}>; |
| ${opts.writeQ?`var<workgroup> pq : array<f32, ${T_WG}>;`:``}`} |
| |
| @compute @workgroup_size(${T_WG}) |
| fn main(@builtin(local_invocation_id) lid : vec3<u32>, |
| @builtin(workgroup_id) wid : vec3<u32>) { |
| let tid = lid.x; |
| let row = tid / RED; |
| let lane = tid % RED; |
| ${normReduce} |
| let oi = wid.x * ROWS + row; // output column in [0,HD) |
| let kb = u.Wk + oi * ${K}u; |
| let vb = u.Wv + oi * ${K}u; |
| ${opts.writeQ?`let qb = u.Wq + oi * ${K}u;`:``} |
| var ka : f32 = 0.0; var va : f32 = 0.0; |
| ${qPart} |
| for (var k : u32 = lane; k < ${K}u; k = k + RED) { |
| let hv = inS[k]; |
| ka = ka + hv * wtv(kb + k); |
| va = va + hv * wtv(vb + k); |
| ${qAccum} |
| } |
| ${reduceWrite} |
| } |
| `; |
| } |
| const buildQonly=()=>gridMatmulT({K:D_MODEL,inBinding:1,outBinding:2,rmsnorm:true,epilogue:"none",hasBias:false}); |
| |
| |
| const ATTN_SINK_WGSL=HEADER+` |
| @group(0) @binding(1) var<storage, read> qbuf : array<f32>; // [HD] |
| @group(0) @binding(2) var<storage, read> kcache : array<f32>; // [cap,HD] |
| @group(0) @binding(3) var<storage, read> vcache : array<f32>; // [cap,HD] |
| @group(0) @binding(4) var<storage, read_write> ctxbuf : array<f32>; // [HD] |
| struct U { per_dim_scale:u32, sink_k:u32, sink_v:u32, cacheLen:u32, }; |
| @group(0) @binding(5) var<uniform> u : U; |
| |
| var<workgroup> logitsW : array<f32, ${KV_CAP + 1}>; |
| var<workgroup> red : array<f32, 256>; |
| var<workgroup> mden : array<f32, 2>; |
| |
| @compute @workgroup_size(256) |
| fn main(@builtin(local_invocation_id) lid : vec3<u32>) { |
| let t = lid.x; |
| let T = u.cacheLen; // number of real keys |
| let nkeys = T + 1u; // + sink |
| let qscale = R_SOFTPLUS_0 * (1.0 / sqrt(f32(UPH))); |
| |
| for (var h : u32 = 0u; h < NUM_HEADS; h = h + 1u) { |
| let qbase = h * UPH; |
| for (var ki : u32 = t; ki < nkeys; ki = ki + 256u) { |
| var acc : f32 = 0.0; |
| if (ki == 0u) { |
| let sbase = u.sink_k + qbase; |
| for (var d : u32 = 0u; d < UPH; d = d + 1u) { |
| acc = acc + qbuf[qbase + d] * wv(sbase + d); |
| } |
| } else { |
| let kk = ki - 1u; |
| let kbase = kk * HD + qbase; |
| for (var d : u32 = 0u; d < UPH; d = d + 1u) { |
| let sc = qscale * softplus(wv(u.per_dim_scale + d)); |
| acc = acc + (qbuf[qbase + d] * sc) * kcache[kbase + d]; |
| } |
| } |
| logitsW[ki] = acc; |
| } |
| workgroupBarrier(); |
| if (t == 0u) { |
| var m : f32 = logitsW[0]; |
| for (var ki : u32 = 1u; ki < nkeys; ki = ki + 1u) { m = max(m, logitsW[ki]); } |
| var s : f32 = 0.0; |
| for (var ki : u32 = 0u; ki < nkeys; ki = ki + 1u) { |
| let e = exp(logitsW[ki] - m); |
| logitsW[ki] = e; s = s + e; |
| } |
| mden[0] = s; |
| } |
| workgroupBarrier(); |
| let denom = mden[0]; |
| for (var d : u32 = t; d < UPH; d = d + 256u) { |
| var acc : f32 = (logitsW[0] / denom) * wv(u.sink_v + qbase + d); |
| for (var kk : u32 = 0u; kk < T; kk = kk + 1u) { |
| let w = logitsW[kk + 1u] / denom; |
| acc = acc + w * vcache[kk * HD + qbase + d]; |
| } |
| ctxbuf[qbase + d] = acc; |
| } |
| workgroupBarrier(); |
| } |
| } |
| `; |
| |
| |
| function gridMatmul2(opts){ |
| const K=opts.K; |
| return HEADER+` |
| @group(0) @binding(${opts.inBinding}) var<storage, read> vin : array<f32>; |
| @group(0) @binding(${opts.outBinding}) var<storage, read_write> vout : array<f32>; |
| struct U { wOff:u32, colBase:u32, normOff:u32, biasOff:u32, }; |
| @group(0) @binding(5) var<uniform> u : U; |
| |
| var<workgroup> inS : array<f32, ${K}>; |
| var<workgroup> part : array<f32, ${WG2}>; |
| |
| @compute @workgroup_size(${WG2}) |
| fn main(@builtin(local_invocation_id) lid : vec3<u32>, |
| @builtin(workgroup_id) wid : vec3<u32>) { |
| let tid = lid.x; |
| let row = tid / LANES; |
| let lane = tid % LANES; |
| for (var i : u32 = tid; i < ${K}u; i = i + ${WG2}u) { inS[i] = vin[i]; } |
| workgroupBarrier(); |
| let outIdx = wid.x * TILE_ROWS + row; |
| let col = u.colBase + outIdx; |
| let wbase = u.wOff + col * ${K}u; |
| var acc : f32 = 0.0; |
| for (var d : u32 = lane; d < ${K}u; d = d + LANES) { |
| acc = acc + inS[d] * wv(wbase + d); |
| } |
| part[tid] = acc; |
| workgroupBarrier(); |
| if (lane == 0u) { |
| var s : f32 = part[row * LANES]; |
| for (var l : u32 = 1u; l < LANES; l = l + 1u) { s = s + part[row * LANES + l]; } |
| vout[outIdx] = s; |
| } |
| } |
| `; |
| } |
| |
| |
| const RESID_WGSL=HEADER+` |
| @group(0) @binding(1) var<storage, read_write> x : array<f32>; |
| @group(0) @binding(2) var<storage, read> y : array<f32>; |
| struct U { post:u32, }; |
| @group(0) @binding(5) var<uniform> u : U; |
| |
| var<workgroup> red : array<f32, 256>; |
| @compute @workgroup_size(256) |
| fn main(@builtin(local_invocation_id) lid : vec3<u32>) { |
| let t = lid.x; |
| var ss : f32 = 0.0; |
| for (var i : u32 = t; i < D_MODEL; i = i + 256u) { let v = y[i]; ss = ss + v*v; } |
| red[t] = ss; |
| workgroupBarrier(); |
| for (var s : u32 = 128u; s > 0u; s = s >> 1u) { |
| if (t < s) { red[t] = red[t] + red[t + s]; } |
| workgroupBarrier(); |
| } |
| let rms = 1.0 / sqrt(red[0] / f32(D_MODEL) + EPS); |
| workgroupBarrier(); |
| for (var i : u32 = t; i < D_MODEL; i = i + 256u) { |
| x[i] = x[i] + (y[i] * rms) * wv(u.post + i); |
| } |
| } |
| `; |
| |
| |
| const EMBED_WGSL=HEADER+` |
| @group(0) @binding(1) var<storage, read> prevBuf : array<u32>; // [12] |
| @group(0) @binding(2) var<storage, read_write> xout : array<f32>; // [1024] |
| struct EU { embeddingOff:u32, embed_scale:f32, }; |
| @group(0) @binding(5) var<uniform> u : EU; |
| |
| @compute @workgroup_size(256) |
| fn main(@builtin(local_invocation_id) lid : vec3<u32>) { |
| let i = lid.x; |
| for (var di : u32 = i; di < D_MODEL; di = di + 256u) { |
| var acc : f32 = 0.0; |
| for (var c : u32 = 0u; c < 12u; c = c + 1u) { |
| let tok = prevBuf[c]; |
| acc = acc + wv(u.embeddingOff + tok * D_MODEL + di); |
| } |
| xout[di] = (acc / 12.0) * u.embed_scale; |
| } |
| } |
| `; |
| |
| |
| |
| const WORD_BYTES=4; |
| const W0_WORDS=W_SPLIT_WORDS, W1_WORDS=STORE_WORDS-W_SPLIT_WORDS; |
| const w0Buf=device.createBuffer({ size:W0_WORDS*WORD_BYTES, usage:GPUBufferUsage.STORAGE|GPUBufferUsage.COPY_DST }); |
| const w1Buf=device.createBuffer({ size:Math.max(4,W1_WORDS*WORD_BYTES), usage:GPUBufferUsage.STORAGE|GPUBufferUsage.COPY_DST }); |
| device.queue.writeBuffer(w0Buf,0,binBytes,0,W0_WORDS*WORD_BYTES); |
| device.queue.writeBuffer(w1Buf,0,binBytes,W0_WORDS*WORD_BYTES,W1_WORDS*WORD_BYTES); |
| |
| |
| const WT_TENSORS=[]; |
| for(let li=0;li<NUM_LAYERS;li++){ |
| WT_TENSORS.push( |
| {src:`L${li}.sa.Wq`,K:D_MODEL,N:HD}, |
| {src:`L${li}.sa.Wk`,K:D_MODEL,N:HD}, |
| {src:`L${li}.sa.Wv`,K:D_MODEL,N:HD}, |
| {src:`L${li}.ca.Wq`,K:D_MODEL,N:HD}, |
| {src:`L${li}.ca.Wk`,K:D_SRC,N:HD}, |
| {src:`L${li}.ca.Wv`,K:D_SRC,N:HD}, |
| {src:`L${li}.ffn.wi`,K:D_MODEL,N:D_FF}, |
| {src:`L${li}.ffn.wo`,K:D_FF,N:D_MODEL}); |
| } |
| const wtOffset={}; let wtFloats=0; |
| for(const t of WT_TENSORS){ wtOffset[t.src]=wtFloats; wtFloats+=t.K*t.N; } |
| const wtOff=n=>wtOffset[n]; |
| |
| const wtF32=new Float32Array(wtFloats); |
| for(const t of WT_TENSORS){ const srcOff=off(t.src), dstOff=wtOffset[t.src], K=t.K, N=t.N; |
| for(let k=0;k<K;k++){ const srcRow=srcOff+k*N; |
| for(let n=0;n<N;n++) wtF32[dstOff+n*K+k]=binF32[srcRow+n]; } } |
| |
| const padWt=wtFloats%2; |
| const wtU16=new Uint16Array(wtFloats+padWt); |
| for(let i=0;i<wtFloats;i++) wtU16[i]=f32ToF16(wtF32[i]); |
| const wtBytes=new Uint8Array(wtU16.buffer); |
| const wtBuf=device.createBuffer({ size:wtBytes.byteLength, usage:GPUBufferUsage.STORAGE|GPUBufferUsage.COPY_DST }); |
| device.queue.writeBuffer(wtBuf,0,wtBytes); |
| |
| |
| const makeStorage=(floats,extra=0)=>device.createBuffer({ |
| size:Math.max(4,floats*4), |
| usage:GPUBufferUsage.STORAGE|GPUBufferUsage.COPY_DST|GPUBufferUsage.COPY_SRC|extra }); |
| const xBuf=makeStorage(D_MODEL); |
| const srcBuf=makeStorage(D_SRC); |
| const qBuf=makeStorage(HD); |
| const ctxBuf=makeStorage(HD); |
| const blkOutBuf=makeStorage(D_MODEL); |
| const ffBuf=makeStorage(D_FF); |
| |
| const selfK=[],selfV=[],crossK=[],crossV=[]; |
| for(let li=0;li<NUM_LAYERS;li++){ |
| selfK.push(makeStorage((KV_CAP+1)*HD)); |
| selfV.push(makeStorage((KV_CAP+1)*HD)); |
| crossK.push(makeStorage((KV_CAP+1)*HD)); |
| crossV.push(makeStorage((KV_CAP+1)*HD)); |
| } |
| const prevBuf=device.createBuffer({ size:12*4, usage:GPUBufferUsage.STORAGE|GPUBufferUsage.COPY_DST }); |
| |
| const makeUniform=bytes=>device.createBuffer({ size:bytes, usage:GPUBufferUsage.UNIFORM|GPUBufferUsage.COPY_DST }); |
| const pipeline=code=>{ const m=device.createShaderModule({code}); return device.createComputePipeline({ layout:"auto", compute:{module:m,entryPoint:"main"} }); }; |
| |
| |
| const pEmbed=pipeline(EMBED_WGSL); |
| const pSelfKV=pipeline(buildKV({K:D_MODEL,writeQ:true,rmsnorm:true})); |
| const pCrossQ=pipeline(buildQonly()); |
| const pCrossKV=pipeline(buildKV({K:D_SRC,writeQ:false,rmsnorm:false})); |
| const pAttn=pipeline(ATTN_SINK_WGSL); |
| const pWo=pipeline(gridMatmul2({K:HD,inBinding:1,outBinding:2})); |
| const pResid=pipeline(RESID_WGSL); |
| const pWi=pipeline(gridMatmulT({K:D_MODEL,inBinding:1,outBinding:2,rmsnorm:true,epilogue:"geluBias",hasBias:true})); |
| const pWoFfn=pipeline(gridMatmulT({K:D_FF,inBinding:1,outBinding:2,rmsnorm:false,epilogue:"none",hasBias:true})); |
| |
| |
| const uEmbed=makeUniform(16); |
| const uSelfKV=[],uCrossQ=[],uCrossKV=[],uAttnSelf=[],uAttnCross=[],uWoSelf=[],uWoCross=[], |
| uResidSAttn=[],uResidCAttn=[],uResidFfn=[],uWi=[],uWoFfn=[]; |
| for(let li=0;li<NUM_LAYERS;li++){ |
| uSelfKV.push(makeUniform(32)); uCrossQ.push(makeUniform(16)); uCrossKV.push(makeUniform(32)); |
| uAttnSelf.push(makeUniform(16)); uAttnCross.push(makeUniform(16)); |
| uWoSelf.push(makeUniform(16)); uWoCross.push(makeUniform(16)); |
| uResidSAttn.push(makeUniform(16)); uResidCAttn.push(makeUniform(16)); uResidFfn.push(makeUniform(16)); |
| uWi.push(makeUniform(16)); uWoFfn.push(makeUniform(16)); |
| } |
| |
| { const e=new ArrayBuffer(16); new Uint32Array(e,0,1)[0]=off("embedding"); new Float32Array(e,4,1)[0]=EMBED_SCALE; |
| device.queue.writeBuffer(uEmbed,0,e); } |
| |
| for(let li=0;li<NUM_LAYERS;li++){ |
| { const a=new Uint32Array(8); a[0]=L(li,"sa","pre_norm"); a[1]=wtOff(`L${li}.sa.Wq`); a[2]=wtOff(`L${li}.sa.Wk`); a[3]=wtOff(`L${li}.sa.Wv`); |
| device.queue.writeBuffer(uSelfKV[li],0,a); } |
| { const a=new Uint32Array(4); a[0]=wtOff(`L${li}.ca.Wq`); a[1]=0; a[2]=L(li,"ca","pre_norm"); a[3]=0; |
| device.queue.writeBuffer(uCrossQ[li],0,a); } |
| { const a=new Uint32Array(8); a[0]=0; a[1]=0; a[2]=wtOff(`L${li}.ca.Wk`); a[3]=wtOff(`L${li}.ca.Wv`); |
| device.queue.writeBuffer(uCrossKV[li],0,a); } |
| { const a=new Uint32Array(4); a[0]=L(li,"sa","per_dim_scale"); a[1]=L(li,"sa","sink_k"); a[2]=L(li,"sa","sink_v"); |
| device.queue.writeBuffer(uAttnSelf[li],0,a); } |
| { const a=new Uint32Array(4); a[0]=L(li,"ca","per_dim_scale"); a[1]=L(li,"ca","sink_k"); a[2]=L(li,"ca","sink_v"); |
| device.queue.writeBuffer(uAttnCross[li],0,a); } |
| { const a=new Uint32Array(4); a[0]=L(li,"sa","Wo"); device.queue.writeBuffer(uWoSelf[li],0,a); } |
| { const a=new Uint32Array(4); a[0]=L(li,"ca","Wo"); device.queue.writeBuffer(uWoCross[li],0,a); } |
| { const a=new Uint32Array(4); a[0]=L(li,"sa","post_norm"); device.queue.writeBuffer(uResidSAttn[li],0,a); } |
| { const a=new Uint32Array(4); a[0]=L(li,"ca","post_norm"); device.queue.writeBuffer(uResidCAttn[li],0,a); } |
| { const a=new Uint32Array(4); a[0]=L(li,"ffn","post_norm"); device.queue.writeBuffer(uResidFfn[li],0,a); } |
| { const a=new Uint32Array(4); a[0]=wtOff(`L${li}.ffn.wi`); a[1]=0; a[2]=L(li,"ffn","pre_norm"); a[3]=L(li,"ffn","bi"); |
| device.queue.writeBuffer(uWi[li],0,a); } |
| { const a=new Uint32Array(4); a[0]=wtOff(`L${li}.ffn.wo`); a[1]=0; a[2]=0; a[3]=L(li,"ffn","bo"); |
| device.queue.writeBuffer(uWoFfn[li],0,a); } |
| } |
| |
| |
| const bgEmbed=device.createBindGroup({ layout:pEmbed.getBindGroupLayout(0), entries:[ |
| {binding:0,resource:{buffer:w0Buf}},{binding:6,resource:{buffer:w1Buf}}, |
| {binding:1,resource:{buffer:prevBuf}},{binding:2,resource:{buffer:xBuf}},{binding:5,resource:{buffer:uEmbed}}] }); |
| const bgSelfKV=Array.from({length:NUM_LAYERS},(_,li)=>device.createBindGroup({ layout:pSelfKV.getBindGroupLayout(0), entries:[ |
| {binding:0,resource:{buffer:w0Buf}},{binding:6,resource:{buffer:w1Buf}},{binding:7,resource:{buffer:wtBuf}}, |
| {binding:1,resource:{buffer:xBuf}},{binding:2,resource:{buffer:qBuf}}, |
| {binding:3,resource:{buffer:selfK[li]}},{binding:4,resource:{buffer:selfV[li]}},{binding:5,resource:{buffer:uSelfKV[li]}}] })); |
| const bgCrossQ=Array.from({length:NUM_LAYERS},(_,li)=>device.createBindGroup({ layout:pCrossQ.getBindGroupLayout(0), entries:[ |
| {binding:0,resource:{buffer:w0Buf}},{binding:6,resource:{buffer:w1Buf}},{binding:7,resource:{buffer:wtBuf}}, |
| {binding:1,resource:{buffer:xBuf}},{binding:2,resource:{buffer:qBuf}},{binding:5,resource:{buffer:uCrossQ[li]}}] })); |
| const bgCrossKV=Array.from({length:NUM_LAYERS},(_,li)=>device.createBindGroup({ layout:pCrossKV.getBindGroupLayout(0), entries:[ |
| |
| {binding:7,resource:{buffer:wtBuf}},{binding:1,resource:{buffer:srcBuf}}, |
| {binding:3,resource:{buffer:crossK[li]}},{binding:4,resource:{buffer:crossV[li]}},{binding:5,resource:{buffer:uCrossKV[li]}}] })); |
| const bgAttnSelf=Array.from({length:NUM_LAYERS},(_,li)=>device.createBindGroup({ layout:pAttn.getBindGroupLayout(0), entries:[ |
| {binding:0,resource:{buffer:w0Buf}},{binding:6,resource:{buffer:w1Buf}}, |
| {binding:1,resource:{buffer:qBuf}},{binding:2,resource:{buffer:selfK[li]}}, |
| {binding:3,resource:{buffer:selfV[li]}},{binding:4,resource:{buffer:ctxBuf}},{binding:5,resource:{buffer:uAttnSelf[li]}}] })); |
| const bgAttnCross=Array.from({length:NUM_LAYERS},(_,li)=>device.createBindGroup({ layout:pAttn.getBindGroupLayout(0), entries:[ |
| {binding:0,resource:{buffer:w0Buf}},{binding:6,resource:{buffer:w1Buf}}, |
| {binding:1,resource:{buffer:qBuf}},{binding:2,resource:{buffer:crossK[li]}}, |
| {binding:3,resource:{buffer:crossV[li]}},{binding:4,resource:{buffer:ctxBuf}},{binding:5,resource:{buffer:uAttnCross[li]}}] })); |
| const bgWoSelf=Array.from({length:NUM_LAYERS},(_,li)=>device.createBindGroup({ layout:pWo.getBindGroupLayout(0), entries:[ |
| {binding:0,resource:{buffer:w0Buf}},{binding:6,resource:{buffer:w1Buf}}, |
| {binding:1,resource:{buffer:ctxBuf}},{binding:2,resource:{buffer:blkOutBuf}},{binding:5,resource:{buffer:uWoSelf[li]}}] })); |
| const bgWoCross=Array.from({length:NUM_LAYERS},(_,li)=>device.createBindGroup({ layout:pWo.getBindGroupLayout(0), entries:[ |
| {binding:0,resource:{buffer:w0Buf}},{binding:6,resource:{buffer:w1Buf}}, |
| {binding:1,resource:{buffer:ctxBuf}},{binding:2,resource:{buffer:blkOutBuf}},{binding:5,resource:{buffer:uWoCross[li]}}] })); |
| const bgResidSAttn=Array.from({length:NUM_LAYERS},(_,li)=>device.createBindGroup({ layout:pResid.getBindGroupLayout(0), entries:[ |
| {binding:0,resource:{buffer:w0Buf}},{binding:6,resource:{buffer:w1Buf}}, |
| {binding:1,resource:{buffer:xBuf}},{binding:2,resource:{buffer:blkOutBuf}},{binding:5,resource:{buffer:uResidSAttn[li]}}] })); |
| const bgResidCAttn=Array.from({length:NUM_LAYERS},(_,li)=>device.createBindGroup({ layout:pResid.getBindGroupLayout(0), entries:[ |
| {binding:0,resource:{buffer:w0Buf}},{binding:6,resource:{buffer:w1Buf}}, |
| {binding:1,resource:{buffer:xBuf}},{binding:2,resource:{buffer:blkOutBuf}},{binding:5,resource:{buffer:uResidCAttn[li]}}] })); |
| const bgWi=Array.from({length:NUM_LAYERS},(_,li)=>device.createBindGroup({ layout:pWi.getBindGroupLayout(0), entries:[ |
| {binding:0,resource:{buffer:w0Buf}},{binding:6,resource:{buffer:w1Buf}},{binding:7,resource:{buffer:wtBuf}}, |
| {binding:1,resource:{buffer:xBuf}},{binding:2,resource:{buffer:ffBuf}},{binding:5,resource:{buffer:uWi[li]}}] })); |
| const bgWoFfn=Array.from({length:NUM_LAYERS},(_,li)=>device.createBindGroup({ layout:pWoFfn.getBindGroupLayout(0), entries:[ |
| {binding:0,resource:{buffer:w0Buf}},{binding:6,resource:{buffer:w1Buf}},{binding:7,resource:{buffer:wtBuf}}, |
| {binding:1,resource:{buffer:ffBuf}},{binding:2,resource:{buffer:blkOutBuf}},{binding:5,resource:{buffer:uWoFfn[li]}}] })); |
| const bgResidFfn=Array.from({length:NUM_LAYERS},(_,li)=>device.createBindGroup({ layout:pResid.getBindGroupLayout(0), entries:[ |
| {binding:0,resource:{buffer:w0Buf}},{binding:6,resource:{buffer:w1Buf}}, |
| {binding:1,resource:{buffer:xBuf}},{binding:2,resource:{buffer:blkOutBuf}},{binding:5,resource:{buffer:uResidFfn[li]}}] })); |
| |
| const WG_KV=tNWGT(HD), WG_Q=tNWGT(HD), WG_WO=tNWG(D_MODEL), WG_WI=tNWGT(D_FF), WG_WOFFN=tNWGT(D_MODEL); |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| let T_LEN=0; |
| |
| function resetKV(){ T_LEN=0; } |
| |
| |
| function setSource(srcF32){ device.queue.writeBuffer(srcBuf,0, srcF32 instanceof Float32Array?srcF32:Float32Array.from(srcF32)); } |
| |
| |
| |
| |
| function runTemporal(prevToks){ |
| const slot=Math.min(T_LEN, KV_CAP); |
| const cacheLen=slot+1; |
| |
| device.queue.writeBuffer(prevBuf,0,new Uint32Array(prevToks)); |
| for(let li=0;li<NUM_LAYERS;li++){ |
| { const a=new Uint32Array(1); a[0]=slot; device.queue.writeBuffer(uSelfKV[li],16,a); device.queue.writeBuffer(uCrossKV[li],16,a); } |
| { const a=new Uint32Array(1); a[0]=cacheLen; device.queue.writeBuffer(uAttnSelf[li],12,a); device.queue.writeBuffer(uAttnCross[li],12,a); } |
| } |
| |
| const enc=device.createCommandEncoder(); |
| const pass=enc.beginComputePass(); |
| pass.setPipeline(pEmbed); pass.setBindGroup(0,bgEmbed); pass.dispatchWorkgroups(1); |
| for(let li=0;li<NUM_LAYERS;li++){ |
| |
| pass.setPipeline(pSelfKV); pass.setBindGroup(0,bgSelfKV[li]); pass.dispatchWorkgroups(WG_KV); |
| pass.setPipeline(pAttn); pass.setBindGroup(0,bgAttnSelf[li]); pass.dispatchWorkgroups(1); |
| pass.setPipeline(pWo); pass.setBindGroup(0,bgWoSelf[li]); pass.dispatchWorkgroups(WG_WO); |
| pass.setPipeline(pResid); pass.setBindGroup(0,bgResidSAttn[li]);pass.dispatchWorkgroups(1); |
| |
| pass.setPipeline(pCrossQ); pass.setBindGroup(0,bgCrossQ[li]); pass.dispatchWorkgroups(WG_Q); |
| pass.setPipeline(pCrossKV); pass.setBindGroup(0,bgCrossKV[li]); pass.dispatchWorkgroups(WG_KV); |
| pass.setPipeline(pAttn); pass.setBindGroup(0,bgAttnCross[li]); pass.dispatchWorkgroups(1); |
| pass.setPipeline(pWo); pass.setBindGroup(0,bgWoCross[li]); pass.dispatchWorkgroups(WG_WO); |
| pass.setPipeline(pResid); pass.setBindGroup(0,bgResidCAttn[li]);pass.dispatchWorkgroups(1); |
| |
| pass.setPipeline(pWi); pass.setBindGroup(0,bgWi[li]); pass.dispatchWorkgroups(WG_WI); |
| pass.setPipeline(pWoFfn); pass.setBindGroup(0,bgWoFfn[li]); pass.dispatchWorkgroups(WG_WOFFN); |
| pass.setPipeline(pResid); pass.setBindGroup(0,bgResidFfn[li]); pass.dispatchWorkgroups(1); |
| } |
| pass.end(); |
| device.queue.submit([enc.finish()]); |
| |
| |
| if(T_LEN<KV_CAP) T_LEN=slot+1; |
| return xBuf; |
| } |
| |
| return { runTemporal, resetKV, setSource, kvLen:()=>T_LEN, outBuf:xBuf }; |
| } |
| |
| |
| |
| |
| |
| |
| |
| const D_IN=1024, D_MODEL=768, D_FF=3072, VOCAB=12294, NUM_CODEBOOKS=12, |
| CODEBOOK_SIZE=1024, NUM_RESERVED=6, NUM_LAYERS=2, NUM_HEADS=6, UPH=128; |
| const HD=NUM_HEADS*UPH, QKV=3*HD, SOFT_CAP=30.0, EPS=1e-6, R_SOFTPLUS_0=1.442695041; |
| const TILE_ROWS=64, LANES=4, WG_SIZE=TILE_ROWS*LANES; |
| const nWG=N=>Math.ceil(N/TILE_ROWS); |
| const RED=32, ROWS=8, T_WG=ROWS*RED; |
| const nWGT=N=>Math.ceil(N/ROWS); |
| let USE_SUBGROUPS=false; |
| |
| const CONSTS=` |
| const D_IN : u32 = ${D_IN}u; |
| const D_MODEL : u32 = ${D_MODEL}u; |
| const D_FF : u32 = ${D_FF}u; |
| const VOCAB : u32 = ${VOCAB}u; |
| const NUM_HEADS : u32 = ${NUM_HEADS}u; |
| const UPH : u32 = ${UPH}u; |
| const HD : u32 = ${HD}u; |
| const QKV : u32 = ${QKV}u; |
| const SOFT_CAP : f32 = ${SOFT_CAP}; |
| const EPS : f32 = ${EPS}; |
| const R_SOFTPLUS_0 : f32 = ${R_SOFTPLUS_0}; |
| const NUM_RESERVED : u32 = ${NUM_RESERVED}u; |
| const CODEBOOK_SIZE : u32 = ${CODEBOOK_SIZE}u; |
| const TILE_ROWS : u32 = ${TILE_ROWS}u; |
| const LANES : u32 = ${LANES}u; |
| const RED : u32 = ${RED}u; |
| const ROWS : u32 = ${ROWS}u; |
| `; |
| function buildHeaders(W_SPLIT){ |
| const WBUFS=` |
| @group(0) @binding(0) var<storage, read> W0 : array<f32>; |
| @group(0) @binding(6) var<storage, read> W1 : array<f32>; |
| const W_SPLIT : u32 = ${W_SPLIT}u; |
| fn wv(i : u32) -> f32 { |
| if (i < W_SPLIT) { return W0[i]; } |
| return W1[i - W_SPLIT]; |
| } |
| fn softplus(x: f32) -> f32 { return log(1.0 + exp(-abs(x))) + max(x, 0.0); } |
| fn gelu_tanh(x: f32) -> f32 { |
| let c : f32 = 0.7978845608028654; |
| return 0.5 * x * (1.0 + tanh(c * (x + 0.044715 * x * x * x))); |
| }`; |
| return CONSTS+WBUFS; |
| } |
| |
| function gridMatmulT(HEADER, opts){ |
| const K=opts.K, sg=USE_SUBGROUPS; |
| const inDecl=`@group(0) @binding(${opts.inBinding}) var<storage, read> vin : array<f32>;`; |
| const outDecl=`@group(0) @binding(${opts.outBinding}) var<storage, read_write> vout : array<f32>;`; |
| const uniformDecl=`struct U { wOff:u32, colBase:u32, normOff:u32, biasOff:u32, }; |
| @group(0) @binding(5) var<uniform> u : U;`; |
| const wtDecl=`@group(0) @binding(7) var<storage, read> WT : array<f32>;`; |
| const normReduce=opts.rmsnorm?` |
| var lss : f32 = 0.0; |
| for (var i : u32 = tid; i < ${K}u; i = i + ${T_WG}u) { let xv = vin[i]; lss = lss + xv * xv; } |
| red[tid] = lss; workgroupBarrier(); |
| for (var s : u32 = ${T_WG/2}u; s > 0u; s = s >> 1u) { if (tid < s) { red[tid] = red[tid] + red[tid + s]; } workgroupBarrier(); } |
| let rms = 1.0 / sqrt(red[0] / f32(${K}u) + EPS); workgroupBarrier(); |
| for (var i : u32 = tid; i < ${K}u; i = i + ${T_WG}u) { inS[i] = (vin[i] * rms) * wv(u.normOff + i); } |
| workgroupBarrier();`:` |
| for (var i : u32 = tid; i < ${K}u; i = i + ${T_WG}u) { inS[i] = vin[i]; } |
| workgroupBarrier();`; |
| let epi=""; |
| if(opts.epilogue==="geluBias") epi=`acc = gelu_tanh(acc + wv(u.biasOff + col));`; |
| else if(opts.epilogue==="softcap"){ const b=opts.hasBias?`acc = acc + wv(u.biasOff + col);`:``; epi=`${b} |
| acc = SOFT_CAP * tanh(acc / SOFT_CAP);`; } |
| else epi=opts.hasBias?`acc = acc + wv(u.biasOff + col);`:``; |
| const reduceWrite=sg?` |
| let total = subgroupAdd(acc); |
| if (lane == 0u) { var accf : f32 = total; ${epi.replace(/\bacc\b/g,"accf")} vout[outIdx] = accf; }`:` |
| part[tid] = acc; workgroupBarrier(); |
| for (var s : u32 = ${RED/2}u; s > 0u; s = s >> 1u) { if (lane < s) { part[tid] = part[tid] + part[tid + s]; } workgroupBarrier(); } |
| if (lane == 0u) { var accf : f32 = part[row * RED]; ${epi.replace(/\bacc\b/g,"accf")} vout[outIdx] = accf; }`; |
| const enableSg=sg?"enable subgroups;\n":""; |
| return enableSg+HEADER+` |
| ${inDecl} |
| ${outDecl} |
| ${uniformDecl} |
| ${wtDecl} |
| var<workgroup> inS : array<f32, ${K}>; |
| ${opts.rmsnorm?`var<workgroup> red : array<f32, ${T_WG}>;`:``} |
| ${sg?"":`var<workgroup> part : array<f32, ${T_WG}>;`} |
| @compute @workgroup_size(${T_WG}) |
| fn main(@builtin(local_invocation_id) lid : vec3<u32>, @builtin(workgroup_id) wid : vec3<u32>) { |
| let tid = lid.x; let row = tid / RED; let lane = tid % RED; |
| ${normReduce} |
| let outIdx = wid.x * ROWS + row; |
| let col = u.colBase + outIdx; |
| let wbase = u.wOff + col * ${K}u; |
| var acc : f32 = 0.0; |
| for (var k : u32 = lane; k < ${K}u; k = k + RED) { acc = acc + inS[k] * WT[wbase + k]; } |
| ${reduceWrite} |
| }`; |
| } |
| |
| function gridMatmul2(HEADER, opts){ |
| const K=opts.K; |
| return HEADER+` |
| @group(0) @binding(${opts.inBinding}) var<storage, read> vin : array<f32>; |
| @group(0) @binding(${opts.outBinding}) var<storage, read_write> vout : array<f32>; |
| struct U { wOff:u32, colBase:u32, normOff:u32, biasOff:u32, }; |
| @group(0) @binding(5) var<uniform> u : U; |
| var<workgroup> inS : array<f32, ${K}>; |
| var<workgroup> part : array<f32, ${WG_SIZE}>; |
| @compute @workgroup_size(${WG_SIZE}) |
| fn main(@builtin(local_invocation_id) lid : vec3<u32>, @builtin(workgroup_id) wid : vec3<u32>) { |
| let tid = lid.x; let row = tid / LANES; let lane = tid % LANES; |
| for (var i : u32 = tid; i < ${K}u; i = i + ${WG_SIZE}u) { inS[i] = vin[i]; } |
| workgroupBarrier(); |
| let outIdx = wid.x * TILE_ROWS + row; |
| let col = u.colBase + outIdx; |
| let wbase = u.wOff + col * ${K}u; |
| var acc : f32 = 0.0; |
| for (var d : u32 = lane; d < ${K}u; d = d + LANES) { acc = acc + inS[d] * wv(wbase + d); } |
| part[tid] = acc; workgroupBarrier(); |
| if (lane == 0u) { var s : f32 = part[row * LANES]; |
| for (var l : u32 = 1u; l < LANES; l = l + 1u) { s = s + part[row * LANES + l]; } |
| vout[outIdx] = s; } |
| }`; |
| } |
| |
| function buildQKV(HEADER){ |
| const sg=USE_SUBGROUPS, enableSg=sg?"enable subgroups;\n":""; |
| const reduceWrite=sg?` |
| let sq = subgroupAdd(qa); let sk = subgroupAdd(ka); let sv = subgroupAdd(va); |
| if (lane == 0u) { qbuf[oi] = sq; let cbase = u.slot * HD; kcache[cbase + oi] = sk; vcache[cbase + oi] = sv; }`:` |
| pq[tid] = qa; pk[tid] = ka; pv[tid] = va; workgroupBarrier(); |
| for (var s : u32 = ${RED/2}u; s > 0u; s = s >> 1u) { |
| if (lane < s) { pq[tid] = pq[tid] + pq[tid + s]; pk[tid] = pk[tid] + pk[tid + s]; pv[tid] = pv[tid] + pv[tid + s]; } |
| workgroupBarrier(); } |
| if (lane == 0u) { qbuf[oi] = pq[row * RED]; let cbase = u.slot * HD; kcache[cbase + oi] = pk[row * RED]; vcache[cbase + oi] = pv[row * RED]; }`; |
| return enableSg+HEADER+` |
| @group(0) @binding(1) var<storage, read> vin : array<f32>; |
| @group(0) @binding(2) var<storage, read_write> qbuf : array<f32>; |
| @group(0) @binding(3) var<storage, read_write> kcache : array<f32>; |
| @group(0) @binding(4) var<storage, read_write> vcache : array<f32>; |
| @group(0) @binding(7) var<storage, read> WT : array<f32>; |
| struct U { normOff:u32, Wq:u32, Wk:u32, Wv:u32, slot:u32, }; |
| @group(0) @binding(5) var<uniform> u : U; |
| var<workgroup> inS : array<f32, ${D_MODEL}>; |
| var<workgroup> red : array<f32, ${T_WG}>; |
| ${sg?"":`var<workgroup> pq : array<f32, ${T_WG}>; |
| var<workgroup> pk : array<f32, ${T_WG}>; |
| var<workgroup> pv : array<f32, ${T_WG}>;`} |
| @compute @workgroup_size(${T_WG}) |
| fn main(@builtin(local_invocation_id) lid : vec3<u32>, @builtin(workgroup_id) wid : vec3<u32>) { |
| let tid = lid.x; let row = tid / RED; let lane = tid % RED; |
| var lss : f32 = 0.0; |
| for (var i : u32 = tid; i < D_MODEL; i = i + ${T_WG}u) { let xv = vin[i]; lss = lss + xv * xv; } |
| red[tid] = lss; workgroupBarrier(); |
| for (var s : u32 = ${T_WG/2}u; s > 0u; s = s >> 1u) { if (tid < s) { red[tid] = red[tid] + red[tid + s]; } workgroupBarrier(); } |
| let rms = 1.0 / sqrt(red[0] / f32(D_MODEL) + EPS); workgroupBarrier(); |
| for (var i : u32 = tid; i < D_MODEL; i = i + ${T_WG}u) { inS[i] = (vin[i] * rms) * wv(u.normOff + i); } |
| workgroupBarrier(); |
| let oi = wid.x * ROWS + row; |
| let qb = u.Wq + oi * HD; let kb = u.Wk + oi * HD; let vb = u.Wv + oi * HD; |
| var qa : f32 = 0.0; var ka : f32 = 0.0; var va : f32 = 0.0; |
| for (var k : u32 = lane; k < D_MODEL; k = k + RED) { let hv = inS[k]; qa = qa + hv * WT[qb + k]; ka = ka + hv * WT[kb + k]; va = va + hv * WT[vb + k]; } |
| ${reduceWrite} |
| }`; |
| } |
| function buildAttnCore(HEADER){ return HEADER+` |
| @group(0) @binding(1) var<storage, read> qbuf : array<f32>; |
| @group(0) @binding(2) var<storage, read> kcache : array<f32>; |
| @group(0) @binding(3) var<storage, read> vcache : array<f32>; |
| @group(0) @binding(4) var<storage, read_write> ctxbuf : array<f32>; |
| struct U { per_dim_scale:u32, cacheLen:u32, }; |
| @group(0) @binding(5) var<uniform> u : U; |
| var<workgroup> logitsW : array<f32, ${NUM_HEADS*NUM_CODEBOOKS}>; |
| var<workgroup> red : array<f32, 2>; |
| @compute @workgroup_size(256) |
| fn main(@builtin(local_invocation_id) lid : vec3<u32>) { |
| let t = lid.x; let T = u.cacheLen; |
| let qscale = R_SOFTPLUS_0 * (1.0 / sqrt(f32(UPH))); |
| for (var h : u32 = 0u; h < NUM_HEADS; h = h + 1u) { |
| for (var k : u32 = t; k < T; k = k + 256u) { |
| var acc : f32 = 0.0; let kbase = k * HD + h * UPH; let qbase = h * UPH; |
| for (var d : u32 = 0u; d < UPH; d = d + 1u) { |
| let sv = qscale * softplus(wv(u.per_dim_scale + d)); |
| acc = acc + (qbuf[qbase + d] * sv) * kcache[kbase + d]; |
| } |
| logitsW[h * T + k] = acc; |
| } |
| } |
| workgroupBarrier(); |
| for (var h : u32 = 0u; h < NUM_HEADS; h = h + 1u) { |
| if (t == 0u) { var m : f32 = logitsW[h * T + 0u]; for (var k : u32 = 1u; k < T; k = k + 1u) { m = max(m, logitsW[h * T + k]); } red[0] = m; } |
| workgroupBarrier(); let m = red[0]; |
| if (t == 0u) { var s : f32 = 0.0; for (var k : u32 = 0u; k < T; k = k + 1u) { let e = exp(logitsW[h * T + k] - m); logitsW[h * T + k] = e; s = s + e; } red[1] = s; } |
| workgroupBarrier(); let denom = red[1]; |
| for (var d : u32 = t; d < UPH; d = d + 256u) { |
| var acc : f32 = 0.0; |
| for (var k : u32 = 0u; k < T; k = k + 1u) { let w = logitsW[h * T + k] / denom; acc = acc + w * vcache[k * HD + h * UPH + d]; } |
| ctxbuf[h * UPH + d] = acc; |
| } |
| workgroupBarrier(); |
| } |
| }`; |
| } |
| function buildResid(HEADER){ return HEADER+` |
| @group(0) @binding(1) var<storage, read_write> x : array<f32>; |
| @group(0) @binding(2) var<storage, read> y : array<f32>; |
| struct U { post:u32, }; |
| @group(0) @binding(5) var<uniform> u : U; |
| var<workgroup> red : array<f32, 256>; |
| @compute @workgroup_size(256) |
| fn main(@builtin(local_invocation_id) lid : vec3<u32>) { |
| let t = lid.x; var ss : f32 = 0.0; |
| for (var i : u32 = t; i < D_MODEL; i = i + 256u) { let v = y[i]; ss = ss + v*v; } |
| red[t] = ss; workgroupBarrier(); |
| for (var s : u32 = 128u; s > 0u; s = s >> 1u) { if (t < s) { red[t] = red[t] + red[t + s]; } workgroupBarrier(); } |
| let rms = 1.0 / sqrt(red[0] / f32(D_MODEL) + EPS); workgroupBarrier(); |
| for (var i : u32 = t; i < D_MODEL; i = i + 256u) { x[i] = x[i] + (y[i] * rms) * wv(u.post + i); } |
| }`; |
| } |
| function buildFinalLN(HEADER){ return HEADER+` |
| @group(0) @binding(1) var<storage, read> x : array<f32>; |
| @group(0) @binding(2) var<storage, read_write> xn : array<f32>; |
| struct U { final_scale:u32, final_bias:u32, }; |
| @group(0) @binding(5) var<uniform> u : U; |
| var<workgroup> red : array<f32, 256>; |
| @compute @workgroup_size(256) |
| fn main(@builtin(local_invocation_id) lid : vec3<u32>) { |
| let t = lid.x; var sm : f32 = 0.0; |
| for (var i : u32 = t; i < D_MODEL; i = i + 256u) { sm = sm + x[i]; } |
| red[t] = sm; workgroupBarrier(); |
| for (var s : u32 = 128u; s > 0u; s = s >> 1u) { if (t < s) { red[t] = red[t] + red[t + s]; } workgroupBarrier(); } |
| let mean = red[0] / f32(D_MODEL); workgroupBarrier(); |
| var vs : f32 = 0.0; |
| for (var i : u32 = t; i < D_MODEL; i = i + 256u) { let dd = x[i] - mean; vs = vs + dd*dd; } |
| red[t] = vs; workgroupBarrier(); |
| for (var s : u32 = 128u; s > 0u; s = s >> 1u) { if (t < s) { red[t] = red[t] + red[t + s]; } workgroupBarrier(); } |
| let rstd = 1.0 / sqrt(red[0] / f32(D_MODEL) + EPS); workgroupBarrier(); |
| for (var i : u32 = t; i < D_MODEL; i = i + 256u) { xn[i] = ((x[i] - mean) * rstd) * wv(u.final_scale + i) + wv(u.final_bias + i); } |
| }`; |
| } |
| |
| function buildSample(){ return CONSTS+` |
| @group(0) @binding(1) var<storage, read> logits1024 : array<f32>; |
| @group(0) @binding(2) var<storage, read> noise : array<f32>; |
| @group(0) @binding(3) var<storage, read_write> outTok : array<u32>; |
| struct U { lo:u32, temperature:f32, slot:u32, noiseBase:u32, }; |
| @group(0) @binding(5) var<uniform> u : U; |
| var<workgroup> bestScore : array<f32, 256>; |
| var<workgroup> bestIdx : array<u32, 256>; |
| @compute @workgroup_size(256) |
| fn main(@builtin(local_invocation_id) lid : vec3<u32>) { |
| let t = lid.x; let lo = u.lo; |
| var bScore : f32 = -3.0e38; var bIdx : u32 = 0xffffffffu; |
| for (var j : u32 = t; j < CODEBOOK_SIZE; j = j + 256u) { |
| var score : f32 = logits1024[j]; |
| let gidx = lo + j; |
| if (u.temperature > 0.0) { |
| var un = noise[u.noiseBase + j]; |
| un = clamp(un, 1.0e-10, 1.0 - 1.0e-7); |
| let g = -log(-log(un)); |
| score = score + g * u.temperature; |
| } |
| if (score > bScore || (score == bScore && gidx < bIdx)) { bScore = score; bIdx = gidx; } |
| } |
| bestScore[t] = bScore; bestIdx[t] = bIdx; workgroupBarrier(); |
| for (var s : u32 = 128u; s > 0u; s = s >> 1u) { |
| if (t < s) { let a = bestScore[t]; let ai = bestIdx[t]; let b = bestScore[t + s]; let bi = bestIdx[t + s]; |
| if (b > a || (b == a && bi < ai)) { bestScore[t] = b; bestIdx[t] = bi; } } |
| workgroupBarrier(); |
| } |
| if (t == 0u) { outTok[u.slot] = bestIdx[0]; } |
| }`; |
| } |
| function buildEmbed(HEADER){ return HEADER+` |
| @group(0) @binding(1) var<storage, read> tokBuf : array<u32>; |
| @group(0) @binding(2) var<storage, read_write> xout : array<f32>; |
| struct EU { embeddingOff : u32, embed_scale : f32, slot : u32 }; |
| @group(0) @binding(3) var<uniform> u : EU; |
| @compute @workgroup_size(256) |
| fn main(@builtin(local_invocation_id) lid : vec3<u32>) { |
| let t = lid.x; let tok = tokBuf[u.slot]; |
| let base = u.embeddingOff + tok * D_IN; |
| for (var i : u32 = t; i < D_IN; i = i + 256u) { xout[i] = wv(base + i) * u.embed_scale; } |
| }`; |
| } |
| |
| |
| async function buildDepth(device, layout, binBytes){ |
| const EMBED_SCALE=layout.embed_scale, TOTAL_FLOATS=layout.total_floats; |
| const W_SPLIT=Math.ceil(TOTAL_FLOATS/2); |
| const off=n=>layout.tensors[n].offset; |
| const layerOffsets=li=>{ const p="l"+li+"_"; return { |
| attn_pre:off(p+"attn_pre"), attn_post:off(p+"attn_post"), |
| Wq:off(p+"Wq"), Wk:off(p+"Wk"), Wv:off(p+"Wv"), |
| per_dim_scale:off(p+"per_dim_scale"), Wo:off(p+"Wo"), |
| ffn_pre:off(p+"ffn_pre"), ffn_post:off(p+"ffn_post"), |
| Wi:off(p+"Wi"), bi:off(p+"bi"), Wo_ffn:off(p+"Wo_ffn"), bo:off(p+"bo") }; }; |
| |
| const makeStorage=(floats,extra=0)=>device.createBuffer({ size:Math.max(4,floats*4), |
| usage:GPUBufferUsage.STORAGE|GPUBufferUsage.COPY_DST|GPUBufferUsage.COPY_SRC|extra }); |
| |
| const W0_BYTES=W_SPLIT*4, W1_BYTES=binBytes.byteLength-W0_BYTES; |
| const w0Buf=device.createBuffer({ size:W0_BYTES, usage:GPUBufferUsage.STORAGE|GPUBufferUsage.COPY_DST }); |
| const w1Buf=device.createBuffer({ size:W1_BYTES, usage:GPUBufferUsage.STORAGE|GPUBufferUsage.COPY_DST }); |
| device.queue.writeBuffer(w0Buf,0,binBytes,0,W0_BYTES); |
| device.queue.writeBuffer(w1Buf,0,binBytes,W0_BYTES,W1_BYTES); |
| |
| const binF32=new Float32Array(binBytes.buffer,binBytes.byteOffset,binBytes.byteLength/4); |
| const WT_TENSORS=[{src:"adapter",K:D_IN,N:D_MODEL}]; |
| for(let li=0;li<NUM_LAYERS;li++){ const p="l"+li+"_"; |
| WT_TENSORS.push({src:p+"Wq",K:D_MODEL,N:HD},{src:p+"Wk",K:D_MODEL,N:HD},{src:p+"Wv",K:D_MODEL,N:HD}, |
| {src:p+"Wi",K:D_MODEL,N:D_FF},{src:p+"Wo_ffn",K:D_FF,N:D_MODEL}); } |
| WT_TENSORS.push({src:"to_logits_w",K:D_MODEL,N:VOCAB}); |
| const wtOffset={}; let wtFloats=0; |
| for(const t of WT_TENSORS){ wtOffset[t.src]=wtFloats; wtFloats+=t.K*t.N; } |
| const wtOff=n=>wtOffset[n]; |
| const wtHost=new Float32Array(wtFloats); |
| for(const t of WT_TENSORS){ const srcOff=off(t.src), dstOff=wtOffset[t.src], K=t.K, N=t.N; |
| for(let k=0;k<K;k++){ const srcRow=srcOff+k*N; for(let n=0;n<N;n++) wtHost[dstOff+n*K+k]=binF32[srcRow+n]; } } |
| const wtBuf=device.createBuffer({ size:wtHost.byteLength, usage:GPUBufferUsage.STORAGE|GPUBufferUsage.COPY_DST }); |
| device.queue.writeBuffer(wtBuf,0,wtHost); |
| |
| const xinBuf=makeStorage(D_IN), x768Buf=makeStorage(D_MODEL), qBuf=makeStorage(HD), ctxBuf=makeStorage(HD), |
| attnOutBuf=makeStorage(D_MODEL), ffBuf=makeStorage(D_FF), ffnOutBuf=makeStorage(D_MODEL), |
| xnBuf=makeStorage(D_MODEL), logits1024Buf=makeStorage(CODEBOOK_SIZE); |
| const kBufs=[makeStorage(NUM_CODEBOOKS*D_MODEL),makeStorage(NUM_CODEBOOKS*D_MODEL)]; |
| const vBufs=[makeStorage(NUM_CODEBOOKS*D_MODEL),makeStorage(NUM_CODEBOOKS*D_MODEL)]; |
| const tokBuf=device.createBuffer({ size:NUM_CODEBOOKS*4, usage:GPUBufferUsage.STORAGE|GPUBufferUsage.COPY_DST|GPUBufferUsage.COPY_SRC }); |
| const tokReadBuf=device.createBuffer({ size:NUM_CODEBOOKS*4, usage:GPUBufferUsage.MAP_READ|GPUBufferUsage.COPY_DST }); |
| const noiseBuf=makeStorage(NUM_CODEBOOKS*CODEBOOK_SIZE); |
| |
| const makeUniform=bytes=>device.createBuffer({ size:bytes, usage:GPUBufferUsage.UNIFORM|GPUBufferUsage.COPY_DST }); |
| const pipeline=code=>{ const m=device.createShaderModule({code}); return device.createComputePipeline({ layout:"auto", compute:{module:m,entryPoint:"main"} }); }; |
| |
| const HEADER=buildHeaders(W_SPLIT); |
| const pAdapter=pipeline(gridMatmulT(HEADER,{K:D_IN,inBinding:1,outBinding:2,rmsnorm:false,epilogue:"none",hasBias:false})); |
| const pQKV=pipeline(buildQKV(HEADER)); |
| const pAttnCore=pipeline(buildAttnCore(HEADER)); |
| const pWo=pipeline(gridMatmul2(HEADER,{K:HD,inBinding:1,outBinding:2})); |
| const pResid=pipeline(buildResid(HEADER)); |
| const pWi=pipeline(gridMatmulT(HEADER,{K:D_MODEL,inBinding:1,outBinding:2,rmsnorm:true,epilogue:"geluBias",hasBias:true})); |
| const pWoFfn=pipeline(gridMatmulT(HEADER,{K:D_FF,inBinding:1,outBinding:2,rmsnorm:false,epilogue:"none",hasBias:true})); |
| const pFinalLN=pipeline(buildFinalLN(HEADER)); |
| const pToLogits=pipeline(gridMatmulT(HEADER,{K:D_MODEL,inBinding:1,outBinding:2,rmsnorm:false,epilogue:"softcap",hasBias:true})); |
| const pSample=pipeline(buildSample()); |
| const pEmbed=pipeline(buildEmbed(HEADER)); |
| |
| const uAdapter=makeUniform(16), uWo=[makeUniform(16),makeUniform(16)], uWi=[makeUniform(16),makeUniform(16)], uWoFfn=[makeUniform(16),makeUniform(16)]; |
| const uToLogits=Array.from({length:NUM_CODEBOOKS},()=>makeUniform(16)); |
| const uQKV=Array.from({length:NUM_CODEBOOKS},()=>[makeUniform(32),makeUniform(32)]); |
| const uAttnCore=Array.from({length:NUM_CODEBOOKS},()=>[makeUniform(16),makeUniform(16)]); |
| const uResidAttn=[makeUniform(16),makeUniform(16)], uResidFfn=[makeUniform(16),makeUniform(16)], uFinalLN=makeUniform(16); |
| const uSample=Array.from({length:NUM_CODEBOOKS},()=>makeUniform(16)); |
| const uEmbed=Array.from({length:NUM_CODEBOOKS},()=>makeUniform(16)); |
| |
| { const a=new Uint32Array(4); a[0]=wtOff("adapter"); a[1]=0; device.queue.writeBuffer(uAdapter,0,a); } |
| for(let li=0;li<NUM_LAYERS;li++){ const Lr=layerOffsets(li), p="l"+li+"_"; |
| for(let q=0;q<NUM_CODEBOOKS;q++){ |
| const a=new Uint32Array(8); a[0]=Lr.attn_pre; a[1]=wtOff(p+"Wq"); a[2]=wtOff(p+"Wk"); a[3]=wtOff(p+"Wv"); a[4]=q; |
| device.queue.writeBuffer(uQKV[q][li],0,a); |
| const c=new Uint32Array(4); c[0]=Lr.per_dim_scale; c[1]=q+1; device.queue.writeBuffer(uAttnCore[q][li],0,c); } |
| { const a=new Uint32Array(4); a[0]=Lr.attn_post; device.queue.writeBuffer(uResidAttn[li],0,a); } |
| { const a=new Uint32Array(4); a[0]=Lr.Wo; a[1]=0; device.queue.writeBuffer(uWo[li],0,a); } |
| { const a=new Uint32Array(4); a[0]=wtOff(p+"Wi"); a[1]=0; a[2]=Lr.ffn_pre; a[3]=Lr.bi; device.queue.writeBuffer(uWi[li],0,a); } |
| { const a=new Uint32Array(4); a[0]=wtOff(p+"Wo_ffn"); a[1]=0; a[2]=0; a[3]=Lr.bo; device.queue.writeBuffer(uWoFfn[li],0,a); } |
| { const a=new Uint32Array(4); a[0]=Lr.ffn_post; device.queue.writeBuffer(uResidFfn[li],0,a); } |
| } |
| { const a=new Uint32Array(4); a[0]=off("final_scale"); a[1]=off("final_bias"); device.queue.writeBuffer(uFinalLN,0,a); } |
| for(let q=0;q<NUM_CODEBOOKS;q++){ const e=new ArrayBuffer(16); |
| new Uint32Array(e,0,1)[0]=off("embedding"); new Float32Array(e,4,1)[0]=EMBED_SCALE; new Uint32Array(e,8,1)[0]=q; |
| device.queue.writeBuffer(uEmbed[q],0,e); } |
| for(let q=0;q<NUM_CODEBOOKS;q++){ const lo=NUM_RESERVED+q*CODEBOOK_SIZE; const a=new Uint32Array(4); |
| a[0]=wtOff("to_logits_w"); a[1]=lo; a[2]=0; a[3]=off("to_logits_b"); device.queue.writeBuffer(uToLogits[q],0,a); } |
| |
| function bgMatmul1(p,inBuf,outBuf,uBuf,usesWv=true){ |
| const entries=[{binding:7,resource:{buffer:wtBuf}},{binding:1,resource:{buffer:inBuf}}, |
| {binding:2,resource:{buffer:outBuf}},{binding:5,resource:{buffer:uBuf}}]; |
| if(usesWv) entries.push({binding:0,resource:{buffer:w0Buf}},{binding:6,resource:{buffer:w1Buf}}); |
| return device.createBindGroup({ layout:p.getBindGroupLayout(0), entries }); |
| } |
| const bgAdapterXin=bgMatmul1(pAdapter,xinBuf,x768Buf,uAdapter,false); |
| const mkAdapterBG=inBuf=>bgMatmul1(pAdapter,inBuf,x768Buf,uAdapter,false); |
| const bgWo=[0,1].map(li=>device.createBindGroup({ layout:pWo.getBindGroupLayout(0), entries:[ |
| {binding:0,resource:{buffer:w0Buf}},{binding:6,resource:{buffer:w1Buf}}, |
| {binding:1,resource:{buffer:ctxBuf}},{binding:2,resource:{buffer:attnOutBuf}},{binding:5,resource:{buffer:uWo[li]}}] })); |
| const bgQKV=Array.from({length:NUM_CODEBOOKS},(_,q)=>[0,1].map(li=>device.createBindGroup({ |
| layout:pQKV.getBindGroupLayout(0), entries:[ |
| {binding:0,resource:{buffer:w0Buf}},{binding:6,resource:{buffer:w1Buf}},{binding:7,resource:{buffer:wtBuf}}, |
| {binding:1,resource:{buffer:x768Buf}},{binding:2,resource:{buffer:qBuf}}, |
| {binding:3,resource:{buffer:kBufs[li]}},{binding:4,resource:{buffer:vBufs[li]}},{binding:5,resource:{buffer:uQKV[q][li]}}] }))); |
| const bgAttnCore=Array.from({length:NUM_CODEBOOKS},(_,q)=>[0,1].map(li=>device.createBindGroup({ |
| layout:pAttnCore.getBindGroupLayout(0), entries:[ |
| {binding:0,resource:{buffer:w0Buf}},{binding:6,resource:{buffer:w1Buf}}, |
| {binding:1,resource:{buffer:qBuf}},{binding:2,resource:{buffer:kBufs[li]}}, |
| {binding:3,resource:{buffer:vBufs[li]}},{binding:4,resource:{buffer:ctxBuf}},{binding:5,resource:{buffer:uAttnCore[q][li]}}] }))); |
| const bgResidAttn=[0,1].map(li=>device.createBindGroup({ layout:pResid.getBindGroupLayout(0), entries:[ |
| {binding:0,resource:{buffer:w0Buf}},{binding:6,resource:{buffer:w1Buf}}, |
| {binding:1,resource:{buffer:x768Buf}},{binding:2,resource:{buffer:attnOutBuf}},{binding:5,resource:{buffer:uResidAttn[li]}}] })); |
| const bgWi=[0,1].map(li=>bgMatmul1(pWi,x768Buf,ffBuf,uWi[li])); |
| const bgWoFfn=[0,1].map(li=>bgMatmul1(pWoFfn,ffBuf,ffnOutBuf,uWoFfn[li])); |
| const bgResidFfn=[0,1].map(li=>device.createBindGroup({ layout:pResid.getBindGroupLayout(0), entries:[ |
| {binding:0,resource:{buffer:w0Buf}},{binding:6,resource:{buffer:w1Buf}}, |
| {binding:1,resource:{buffer:x768Buf}},{binding:2,resource:{buffer:ffnOutBuf}},{binding:5,resource:{buffer:uResidFfn[li]}}] })); |
| const bgFinalLN=device.createBindGroup({ layout:pFinalLN.getBindGroupLayout(0), entries:[ |
| {binding:0,resource:{buffer:w0Buf}},{binding:6,resource:{buffer:w1Buf}}, |
| {binding:1,resource:{buffer:x768Buf}},{binding:2,resource:{buffer:xnBuf}},{binding:5,resource:{buffer:uFinalLN}}] }); |
| const bgToLogits=Array.from({length:NUM_CODEBOOKS},(_,q)=>bgMatmul1(pToLogits,xnBuf,logits1024Buf,uToLogits[q])); |
| const bgSample=Array.from({length:NUM_CODEBOOKS},(_,q)=>device.createBindGroup({ layout:pSample.getBindGroupLayout(0), entries:[ |
| {binding:1,resource:{buffer:logits1024Buf}},{binding:2,resource:{buffer:noiseBuf}}, |
| {binding:3,resource:{buffer:tokBuf}},{binding:5,resource:{buffer:uSample[q]}}] })); |
| const bgEmbed=Array.from({length:NUM_CODEBOOKS},(_,q)=>device.createBindGroup({ layout:pEmbed.getBindGroupLayout(0), entries:[ |
| {binding:0,resource:{buffer:w0Buf}},{binding:6,resource:{buffer:w1Buf}}, |
| {binding:1,resource:{buffer:tokBuf}},{binding:2,resource:{buffer:xinBuf}},{binding:3,resource:{buffer:uEmbed[q]}}] })); |
| |
| const WG_ADAPTER=nWGT(D_MODEL), WG_QKV=nWGT(HD), WG_WO=nWG(D_MODEL), |
| WG_WI=nWGT(D_FF), WG_WOFFN=nWGT(D_MODEL), WG_TOLOGITS=nWGT(CODEBOOK_SIZE); |
| |
| const adapterBGCache=new WeakMap(); |
| |
| |
| |
| |
| async function runDepth(temporalOutBuf, noiseFlat, temperature){ |
| device.queue.writeBuffer(noiseBuf,0,noiseFlat); |
| for(let q=0;q<NUM_CODEBOOKS;q++){ const lo=NUM_RESERVED+q*CODEBOOK_SIZE; const buf=new ArrayBuffer(16); |
| new Uint32Array(buf,0,1)[0]=lo; new Float32Array(buf,4,1)[0]=temperature; |
| new Uint32Array(buf,8,1)[0]=q; new Uint32Array(buf,12,1)[0]=q*CODEBOOK_SIZE; |
| device.queue.writeBuffer(uSample[q],0,buf); } |
| |
| let bgAdapter0=adapterBGCache.get(temporalOutBuf); |
| if(!bgAdapter0){ bgAdapter0=mkAdapterBG(temporalOutBuf); adapterBGCache.set(temporalOutBuf,bgAdapter0); } |
| |
| for(let q=0;q<NUM_CODEBOOKS;q++){ |
| const enc=device.createCommandEncoder(); const pass=enc.beginComputePass(); |
| pass.setPipeline(pAdapter); |
| pass.setBindGroup(0, q===0?bgAdapter0:bgAdapterXin); |
| pass.dispatchWorkgroups(WG_ADAPTER); |
| for(let li=0;li<NUM_LAYERS;li++){ |
| pass.setPipeline(pQKV); pass.setBindGroup(0,bgQKV[q][li]); pass.dispatchWorkgroups(WG_QKV); |
| pass.setPipeline(pAttnCore); pass.setBindGroup(0,bgAttnCore[q][li]); pass.dispatchWorkgroups(1); |
| pass.setPipeline(pWo); pass.setBindGroup(0,bgWo[li]); pass.dispatchWorkgroups(WG_WO); |
| pass.setPipeline(pResid); pass.setBindGroup(0,bgResidAttn[li]); pass.dispatchWorkgroups(1); |
| pass.setPipeline(pWi); pass.setBindGroup(0,bgWi[li]); pass.dispatchWorkgroups(WG_WI); |
| pass.setPipeline(pWoFfn); pass.setBindGroup(0,bgWoFfn[li]); pass.dispatchWorkgroups(WG_WOFFN); |
| pass.setPipeline(pResid); pass.setBindGroup(0,bgResidFfn[li]); pass.dispatchWorkgroups(1); |
| } |
| pass.setPipeline(pFinalLN); pass.setBindGroup(0,bgFinalLN); pass.dispatchWorkgroups(1); |
| pass.setPipeline(pToLogits); pass.setBindGroup(0,bgToLogits[q]); pass.dispatchWorkgroups(WG_TOLOGITS); |
| pass.setPipeline(pSample); pass.setBindGroup(0,bgSample[q]); pass.dispatchWorkgroups(1); |
| if(q<NUM_CODEBOOKS-1){ pass.setPipeline(pEmbed); pass.setBindGroup(0,bgEmbed[q]); pass.dispatchWorkgroups(1); } |
| pass.end(); device.queue.submit([enc.finish()]); |
| } |
| const enc2=device.createCommandEncoder(); |
| enc2.copyBufferToBuffer(tokBuf,0,tokReadBuf,0,NUM_CODEBOOKS*4); |
| device.queue.submit([enc2.finish()]); |
| await tokReadBuf.mapAsync(GPUMapMode.READ); |
| const toksU=new Uint32Array(tokReadBuf.getMappedRange().slice(0)); tokReadBuf.unmap(); |
| const toks=[], codes=[]; |
| for(let q=0;q<NUM_CODEBOOKS;q++){ const tok=Number(toksU[q]); toks.push(tok); |
| codes.push(((tok-NUM_RESERVED)%CODEBOOK_SIZE+CODEBOOK_SIZE)%CODEBOOK_SIZE); } |
| return {codes, toks}; |
| } |
| return { runDepth }; |
| } |
| |
| |
| const NR=()=>cfg.num_reserved_tokens, CB=()=>cfg.codebook_size, Q=()=>cfg.num_codebooks, D=()=>cfg.model_dims; |
| const i64=(arr,dims)=>new ort.Tensor("int64",BigInt64Array.from(arr.map(x=>BigInt(x))),dims); |
| const f32=(arr,dims)=>new ort.Tensor("float32",arr instanceof Float32Array?arr:Float32Array.from(arr),dims); |
| |
| |
| async function encodeStyle(tokens){ |
| const off=cfg.conditioning.cond_offset; |
| const vals=[...tokens, ...Array(128).fill(-1), -1, 20,10,8].map(v=>v+off); |
| const {source}=await S.encoder.run({cond:i64(vals,[1,1,vals.length])}); |
| const data = source.data ?? await source.getData(); |
| const srcF32 = data instanceof Float32Array ? Float32Array.from(data) : Float32Array.from(data); |
| try{ source.dispose && source.dispose(); }catch(e){} |
| return srcF32; |
| } |
| |
| |
| |
| let prevToks=null, sourceF32=null, history=[], pendingStyle=null; |
| function dispose(t){ try{ t && t.dispose && t.dispose(); }catch(e){} } |
| function reset(){ |
| |
| prevToks=Array(Q()).fill(0); |
| history=[]; |
| if(TEMPORAL) TEMPORAL.resetKV(); |
| if(DSJ){ if(DSTATE) DSJ.keys.forEach((k,i)=>dispose(DSTATE["s_"+i])); |
| DSTATE={}; DSJ.keys.forEach((k,i)=>{ const sh=DSJ.shapes[k]; DSTATE["s_"+i]=f32(new Float32Array(sh.reduce((a,b)=>a*b,1)),sh); }); } |
| } |
| |
| |
| let DEC_MS=0, DEPTH_MS=0, TEMP_MS=0; |
| async function genFrame(temp){ |
| |
| |
| const t0=performance.now(); |
| const outBuf=TEMPORAL.runTemporal(prevToks); |
| TEMP_MS=performance.now()-t0; |
| |
| const noise=Float32Array.from({length:Q()*CODEBOOK_SIZE},()=>Math.random()); |
| const t1=performance.now(); |
| const {codes,toks}=await DEPTH.runDepth(outBuf, noise, temp); |
| DEPTH_MS=performance.now()-t1; |
| prevToks=toks; |
| return codes; |
| } |
| |
| |
| function codesToEmb(frames){ const F=frames.length, dd=256, out=new Float32Array(F*dd); |
| for(let t=0;t<F;t++) for(let q=0;q<12;q++){ const c=frames[t][q], o=(q*1024+c)*dd; |
| for(let i=0;i<dd;i++) out[t*dd+i]+=QUANT[o+i]; } return f32(out,[1,F,dd]); } |
| |
| async function streamDecode(frames){ |
| const emb=codesToEmb(frames); |
| const feed={embeddings:emb}; DSJ.keys.forEach((k,i)=>feed["s_"+i]=DSTATE["s_"+i]); |
| const r=await S.decoder.run(feed); |
| DSJ.keys.forEach((k,i)=>{ dispose(DSTATE["s_"+i]); DSTATE["s_"+i]=r["ns_"+i]; }); |
| dispose(emb); |
| return r.waveform; |
| } |
| |
| |
| $("load").onclick=async()=>{ |
| $("load").disabled=true; $("prog").style.display="block"; setBar(0); |
| $("stat").textContent=`loading ${sz} · webgpu · custom WGSL temporal + WGSL depth…`; |
| const base=REPO; |
| try{ |
| cfg=await fetchJSON(base+"/config.json"); |
| presets=await fetchJSON(base+"/presets.json"); |
| DSJ=await fetchJSON(DEC_STATE_URL); |
| $("prompt").innerHTML=Object.keys(presets).map(p=>`<option>${p}</option>`).join(""); |
| const ep=["webgpu"]; |
| HIFLUSH={extra:{"ep.webgpuexecutionprovider.maxPendingDispatches":"100000"}}; |
| |
| const g={encoder:"encoder",decoder:"spectrostream_decoder"}; |
| const extMap=cfg.external_data||{}; |
| const items=[]; |
| for(const k in g){ |
| const fn = k==="decoder" ? "spectrostream_stream.onnx" : `${g[k]}.q4.onnx`; |
| items.push({id:"g_"+k, url: k==="decoder" ? DEC_URL : `${base}/onnx/${fn}`}); |
| const df=extMap[fn]; if(df) items.push({id:"d_"+k, path:df, url:`${base}/onnx/${df}`}); } |
| items.push({id:"emb", url:base+"/onnx/token_embedding_fp16.bin"}); |
| items.push({id:"quant", url:base+"/onnx/quantizer_fp16.bin"}); |
| |
| items.push({id:"temp_layout", url:TEMP_BASE+"temporal_layout.json", json:true}); |
| items.push({id:"temp_bin", url:TEMP_BASE+"temporal.bin"}); |
| |
| items.push({id:"depth_layout", url:DEPTH_BASE+"depth_layout.json", json:true}); |
| items.push({id:"depth_bin", url:DEPTH_BASE+"depth.bin"}); |
| const N=items.length, bufs={}; |
| for(let i=0;i<N;i++){ const it=items[i], name=it.url.split("/").pop(); |
| if(it.json){ bufs[it.id]=await fetchJSON(it.url); setBar((i+1)/N); continue; } |
| bufs[it.id]=await fetchBuf(it.url, f=>{ setBar((i+f)/N); $("stat").textContent=`downloading ${name} — ${(f*100).toFixed(0)}% (${i+1}/${N})`; }); } |
| setBar(1); $("stat").textContent="compiling ORT graphs…"; |
| for(const k in g){ const opts={executionProviders:ep}, d=items.find(x=>x.id==="d_"+k); |
| if(d) opts.externalData=[{path:d.path, data:bufs["d_"+k]}]; |
| if(k==="decoder"){ opts.preferredOutputLocation={}; DSJ.keys.forEach((_,i)=>opts.preferredOutputLocation["ns_"+i]="gpu-buffer"); } |
| S[k]=await ort.InferenceSession.create(bufs["g_"+k], opts); } |
| EMB=f16decode(bufs.emb); QUANT=f16decode(bufs.quant); |
| |
| |
| const ortDevice = ort.env.webgpu.device; |
| if(!ortDevice) throw new Error("ort.env.webgpu.device is null — WebGPU EP not initialized"); |
| $("stat").textContent="building WGSL temporal on ORT device…"; |
| TEMPORAL = await buildTemporal(ortDevice, bufs.temp_layout, bufs.temp_bin); |
| $("stat").textContent="building WGSL depth on ORT device…"; |
| DEPTH = await buildDepth(ortDevice, bufs.depth_layout, bufs.depth_bin); |
| |
| |
| $("stat").textContent="warming up (compiling shaders)…"; setBar(0); |
| try{ reset(); sourceF32=await encodeStyle(presets[$("prompt").value]); TEMPORAL.setSource(sourceF32); |
| const W=(cfg.temporal&&cfg.temporal.max_past?cfg.temporal.max_past:T_KV_CAP)+4; |
| for(let w=0;w<W;w++){ const wt=performance.now(); const c=await genFrame(1.1); history.push(c); const ms=performance.now()-wt; |
| setBar((w+1)/W); $("stat").textContent=`warming up ${w+1}/${W} · ${ms.toFixed(0)} ms/frame (temporal ${TEMP_MS.toFixed(1)} · depth ${DEPTH_MS.toFixed(1)} ms)`; |
| await new Promise(r=>setTimeout(r)); } |
| if(history.length>=10){ const w=await streamDecode(history.slice(-10)); dispose(w); } |
| reset(); }catch(e){ console.warn("warmup:",e); } |
| $("prog").style.display="none"; $("play").disabled=false; |
| $("stat").innerHTML=`ready · <code>webgpu</code> · <code>WGSL temporal + WGSL depth</code> — warmed, press Play`; |
| }catch(e){ $("prog").style.display="none"; $("stat").textContent="load failed: "+e.message; $("load").disabled=false; console.error(e); } |
| }; |
| |
| |
| let actx=null,analyser=null,nextT=0,playing=false,freq=null; |
| function playSamples(data,start,end){ const n=end-start; if(n<1)return; |
| const buf=actx.createBuffer(2,n,48000), L=buf.getChannelData(0), R=buf.getChannelData(1); |
| for(let i=0;i<n;i++){ L[i]=data[(start+i)*2]; R[i]=data[(start+i)*2+1]; } |
| const s=actx.createBufferSource(); s.buffer=buf; s.connect(analyser); |
| if(nextT<actx.currentTime+0.05)nextT=actx.currentTime+0.15; s.start(nextT); nextT+=n/48000; } |
| function viz(){ const c=$("viz"),g=c.getContext("2d"),w=c.width,h=c.height; g.clearRect(0,0,w,h); |
| if(analyser){ analyser.getByteFrequencyData(freq); const bw=w/48; |
| for(let i=0;i<48;i++){ const v=freq[i*2]/255; g.fillStyle=`hsl(${70-v*40},85%,${45+v*15}%)`; g.fillRect(i*bw,h-v*h,bw-1,v*h);} } |
| requestAnimationFrame(viz); } viz(); |
| |
| $("play").onclick=async()=>{ |
| actx=new AudioContext({sampleRate:48000}); await actx.resume(); nextT=actx.currentTime+0.2; |
| analyser=actx.createAnalyser(); analyser.fftSize=128; analyser.connect(actx.destination); freq=new Uint8Array(analyser.frequencyBinCount); |
| $("play").disabled=true; $("stop").disabled=false; playing=true; |
| reset(); sourceF32=await encodeStyle(presets[$("prompt").value]); TEMPORAL.setSource(sourceF32); |
| const CHUNK=10; let done=0, firstChunk=true; const t0=performance.now(); |
| |
| let aT=0,aD=0,aC=0,nT=0; |
| while(playing){ |
| if(pendingStyle){ const old=sourceF32; sourceF32=await encodeStyle(presets[pendingStyle]); TEMPORAL.setSource(sourceF32); pendingStyle=null; } |
| const nf=[]; |
| for(let i=0;i<CHUNK&&playing;i++){ const c=await genFrame(1.1); history.push(c); nf.push(c); done++; |
| aT+=TEMP_MS; aD+=DEPTH_MS; nT++; } |
| if(nf.length===CHUNK){ |
| const dt=performance.now(); |
| const wav=await streamDecode(nf); |
| const data=wav.data ?? await wav.getData(); DEC_MS=performance.now()-dt; |
| aC+=DEC_MS; |
| const start=firstChunk?1920:0; firstChunk=false; |
| playSamples(data, start, wav.dims[1]); dispose(wav); |
| } |
| const sec=done/cfg.frames_per_second, rt=sec/((performance.now()-t0)/1000); |
| $("rt").textContent=`${rt.toFixed(2)}× realtime`; |
| $("stat").textContent=`generating · ${sec.toFixed(1)}s · ${rt.toFixed(2)}× realtime`; |
| |
| const chunks=Math.max(1,Math.round(nT/CHUNK)); |
| $("brk").textContent=`temporal ${(aT/Math.max(1,nT)).toFixed(1)} ms · depth ${(aD/Math.max(1,nT)).toFixed(1)} ms · decode ${(aC/Math.max(1,chunks)).toFixed(1)} ms/chunk`; |
| } |
| }; |
| $("stop").onclick=()=>{ playing=false; if(actx)actx.suspend(); $("play").disabled=false; $("stop").disabled=true; $("stat").textContent="stopped"; }; |
| $("prompt").onchange=()=>{ if(playing) pendingStyle=$("prompt").value; }; |
| </script> |
| </body> |
| </html> |
|
|