| <!DOCTYPE html> |
| <html lang="en"> |
| <head> |
| <meta charset="utf-8"/> |
| <meta name="viewport" content="width=device-width, initial-scale=1"/> |
| <title>Magenta RT · Live (ORT 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; } |
| 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">WGSL depth on your GPU</span></h1> |
| <div class="sub">ORT temporal → <b>custom WGSL depth</b> (GPU-resident, shared ORT device) → ORT streaming decoder. Hear the depth speedup on Metal.</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 temporal + decoder + 148 MB WGSL depth weights, cached after first run).</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", dtype = "q4"; |
| 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/"; |
| let cfg=null, S={}, EMB=null, QUANT=null, presets={}; |
| let DSJ=null, DSTATE=null, HIFLUSH=undefined; |
| let DEPTH=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-wgsldepth-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 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]; // compact [Q*1024] noise |
| 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 L=layerOffsets(li), p="l"+li+"_"; |
| for(let q=0;q<NUM_CODEBOOKS;q++){ |
| const a=new Uint32Array(8); a[0]=L.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]=L.per_dim_scale; c[1]=q+1; device.queue.writeBuffer(uAttnCore[q][li],0,c); } |
| { const a=new Uint32Array(4); a[0]=L.attn_post; device.queue.writeBuffer(uResidAttn[li],0,a); } |
| { const a=new Uint32Array(4); a[0]=L.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]=L.ffn_pre; a[3]=L.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]=L.bo; device.queue.writeBuffer(uWoFfn[li],0,a); } |
| { const a=new Uint32Array(4); a[0]=L.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); |
| function emptyKV(L,nh,uph){ return f32(new Float32Array(0),[L,1,0,nh,uph]); } |
| |
| |
| 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])}); |
| return source; |
| } |
| |
| |
| let KV=null, prev=null, source=null, history=[], pendingStyle=null; |
| function dispose(t){ try{ t && t.dispose && t.dispose(); }catch(e){} } |
| function reset(){ const t=cfg.temporal; |
| KV={sk:emptyKV(t.num_layers,t.num_heads,t.dim_per_head), sv:emptyKV(t.num_layers,t.num_heads,t.dim_per_head), |
| ck:emptyKV(t.num_layers,t.num_heads,t.dim_per_head), cv:emptyKV(t.num_layers,t.num_heads,t.dim_per_head)}; |
| prev=i64(Array(Q()).fill(0),[1,1,Q()]); history=[]; |
| 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; |
| async function genFrame(temp){ |
| |
| const r=await S.temporal.run({prev,self_k:KV.sk,self_v:KV.sv,cross_k:KV.ck,cross_v:KV.cv,source}, HIFLUSH); |
| dispose(KV.sk);dispose(KV.sv);dispose(KV.ck);dispose(KV.cv);dispose(prev); |
| KV={sk:r.new_self_k,sv:r.new_self_v,ck:r.new_cross_k,cv:r.new_cross_v}; |
| |
| const outBuf=r.out.gpuBuffer; |
| |
| const noise=Float32Array.from({length:Q()*CODEBOOK_SIZE},()=>Math.random()); |
| const t0=performance.now(); |
| const {codes,toks}=await DEPTH.runDepth(outBuf, noise, temp); |
| DEPTH_MS=performance.now()-t0; |
| dispose(r.out); |
| prev=i64(toks,[1,1,Q()]); |
| 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 · q4 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",temporal:"temporal_step",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:"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==="temporal") opts.preferredOutputLocation={new_self_k:"gpu-buffer",new_self_v:"gpu-buffer",new_cross_k:"gpu-buffer",new_cross_v:"gpu-buffer",out:"gpu-buffer"}; |
| 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); |
| |
| |
| $("stat").textContent="building WGSL depth on ORT device…"; |
| const ortDevice = ort.env.webgpu.device; |
| if(!ortDevice) throw new Error("ort.env.webgpu.device is null — WebGPU EP not initialized"); |
| DEPTH = await buildDepth(ortDevice, bufs.depth_layout, bufs.depth_bin); |
| |
| |
| $("stat").textContent="warming up (compiling shaders)…"; setBar(0); |
| try{ reset(); source=await encodeStyle(presets[$("prompt").value]); |
| const W=cfg.temporal.max_past+4; |
| for(let w=0;w<W;w++){ const wt=performance.now(); await genFrame(1.1); const ms=performance.now()-wt; |
| setBar((w+1)/W); $("stat").textContent=`warming up ${w+1}/${W} · ${ms.toFixed(0)} ms/frame (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>q4 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(); source=await encodeStyle(presets[$("prompt").value]); |
| const CHUNK=10; let done=0, firstChunk=true; const t0=performance.now(); |
| while(playing){ |
| if(pendingStyle){ const old=source; source=await encodeStyle(presets[pendingStyle]); pendingStyle=null; dispose(old); } |
| const nf=[]; |
| for(let i=0;i<CHUNK&&playing;i++){ const c=await genFrame(1.1); history.push(c); nf.push(c); done++; } |
| 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; |
| 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 · depth ${DEPTH_MS.toFixed(1)} ms/frame · decode ${DEC_MS.toFixed(0)} 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> |
|
|