agate-webgpu / js /app.js
Stefatorus
Claude Opus 5.5
Faster WebGPU sampling: fp16 compute, GPU-resident steps, non-blocking thinker preview
146ba97
Raw History Blame Contribute Delete
17.7 kB
import { Agate, VERSIONS, webgpuStatus, clearCache, hasCachedModel } from "./agate.js";
import { addPngText, readWatermark } from "./marking.js";
import { prepare as preparePrompt, countVector } from "./prompt.js";
import { ThinkerView } from "./thinker.js";
// Model location: 001 from window.AGATE_MODEL_BASE (js/config.js, default ./models/), 002 / 003 from their model
// repos (agate.js VERSIONS); ?models=<url> serves every version from one local-layout base instead.
const params = new URLSearchParams(location.search);
const BASE = window.AGATE_MODEL_BASE || "./models/"; // 001's files (002 / 003 come from their model repos)
const MODELS_OVERRIDE = params.get("models"); // every version from <url>/{,002/,003/} (local testing)
const FORCE_EP = params.get("ep"); // "wasm" forces the CPU path
const PARITY = params.has("parity");
const GOLDEN = params.has("golden"); // prompt-pipeline golden test (003)
const VARIANT = params.get("variant"); // "full16": experimental fp16-compute generator (001)
const store = {
get(k) { try { return localStorage.getItem(k); } catch { return null; } },
set(k, v) { try { localStorage.setItem(k, v); } catch { /* storage blocked */ } },
};
const DEFAULT_VERSION = "003";
let version = [params.get("v"), store.get("agate-version"), DEFAULT_VERSION].find((v) => v && VERSIONS[v]);
let resolution = Number(params.get("res") || store.get("agate-res-003") || 512);
if (!VERSIONS["003"].res.includes(resolution)) resolution = 512;
const EXAMPLES = {
common: [
"a green teapot and a red cup on a table",
"a minimalist logo of a fox head, orange, flat design, white background",
"a dog sitting to the left of a cat",
"a red cube on top of a blue sphere",
],
"003": [
'a shop sign that says "OPEN", three red apples, no people',
"a portrait of an old fisherman at golden hour, detailed",
'a coffee cup with the text "Good Morning" on it',
],
};
const $ = (id) => document.getElementById(id);
const ui = {
prompt: $("prompt"), negative: $("negative"), seed: $("seed"), steps: $("steps"), cfg: $("cfg"), go: $("go"), dice: $("dice"),
fill: $("bar-fill"), status: $("status"), statusR: $("status-r"), canvas: $("canvas"), frame: $("frame"),
placeholder: $("placeholder"), timings: $("timings"), save: $("save"), crisp: $("crisp"), backend: $("backend"),
versions: $("versions"), res: $("res"), resRow: $("res-row"), prepared: $("prepared"), aiNote: $("ai-note"),
showThinker: $("show-thinker"), thinker: $("thinker"), thinkerStep: $("thinker-step"),
};
const thinkerView = new ThinkerView($("plan-canvas"), $("pred-canvas"));
ui.showThinker.checked = store.get("agate-thinker") !== "0";
ui.showThinker.onchange = () => { store.set("agate-thinker", ui.showThinker.checked ? "1" : "0"); if (!ui.showThinker.checked) ui.thinker.hidden = true; };
let agate = null, ep = FORCE_EP || "webgpu", busy = false, stop = false, gpuOk = true, last = null;
window.__agate = { state: "idle", version }; // read by the automated browser test
const res = () => (version === "003" ? resolution : 256);
function renderExamples() {
const box = $("examples");
box.querySelectorAll(".chip").forEach((c) => c.remove());
for (const p of [...(EXAMPLES[version] || []), ...EXAMPLES.common].slice(0, 6)) {
const b = document.createElement("button");
b.className = "chip"; b.textContent = p; b.type = "button";
b.onclick = () => { ui.prompt.value = p; };
box.appendChild(b);
}
}
function seg(el, items, current, onPick) {
el.innerHTML = "";
for (const it of items) {
const b = document.createElement("button");
b.type = "button"; b.setAttribute("role", "radio"); b.setAttribute("aria-checked", String(it.value === current));
b.dataset.value = it.value;
b.innerHTML = `${it.label}${it.sub ? `<small>${it.sub}</small>` : ""}`;
b.onclick = () => { if (!busy && it.value !== current) onPick(it.value); };
el.appendChild(b);
}
}
function renderVersion() {
const spec = VERSIONS[version];
seg(ui.versions, Object.keys(VERSIONS).map((v) => ({ value: v, label: `Preview ${v}`, sub: `${Math.max(...VERSIONS[v].res)} px${v === DEFAULT_VERSION ? " · newest" : ""}` })),
version, pickVersion);
$("ver-note").textContent = spec.note;
$("mast-version").textContent = `Preview ${version}`;
$("mast-spec").textContent = `0.26B · text-to-image · ${spec.res.join(" / ")} px`;
$("ai-model").textContent = `Agate Preview ${version}`;
$("neg-hint").textContent = version === "003" ? "(optional; “no X” in the prompt is added automatically)" : "(optional negative prompt)";
ui.resRow.hidden = spec.res.length < 2;
if (spec.res.length > 1) seg(ui.res, spec.res.map((r) => ({ value: r, label: `${r} × ${r}`, sub: r === 512 ? "native" : "faster" })), resolution, pickRes);
$("out-label").textContent = `Output · ${res()} × ${res()}`;
renderExamples();
document.title = `Agate Preview ${version} — in your browser`;
}
async function pickVersion(v) {
version = v; store.set("agate-version", v);
window.__agate.version = v;
if (agate) { await agate.release(); agate = null; }
clearOutput();
renderVersion();
if (!gpuOk && !FORCE_EP) return;
ui.go.textContent = "Load model"; ui.go.disabled = false;
setProgress(0, "Idle", "");
if (await hasCachedModel(v)) await load();
}
async function pickRes(r) {
resolution = r; store.set("agate-res-003", String(r));
renderVersion();
if (agate?.sessions && agate.mr) {
busy = true; ui.go.disabled = true;
setProgress(1, `Preparing the ${r} px generator`, "");
try { const ms = await agate.setResolution(r); setProgress(0, "Ready", `${r} px · ${(ms / 1000).toFixed(1)} s`); }
catch (e) { console.error(e); setProgress(0, "Switch failed", String(e.message || e)); }
busy = false; ui.go.disabled = false;
}
}
ui.steps.oninput = () => { $("steps-v").textContent = ui.steps.value; $("fast").checked = Number(ui.steps.value) === 25; };
$("fast").onchange = () => { ui.steps.value = $("fast").checked ? 25 : 50; $("steps-v").textContent = ui.steps.value; };
ui.cfg.oninput = () => { $("cfg-v").textContent = Number(ui.cfg.value).toFixed(1); };
ui.dice.onclick = () => { ui.seed.value = Math.floor(Math.random() * 2 ** 31); };
ui.crisp.checked = store.get("agate-crisp") === "1";
const applyCrisp = () => { ui.frame.classList.toggle("crisp", ui.crisp.checked); store.set("agate-crisp", ui.crisp.checked ? "1" : "0"); };
ui.crisp.onchange = applyCrisp; applyCrisp();
const mb = (b) => (b / 1e6).toFixed(0);
function setProgress(frac, left, right = "") {
ui.fill.style.width = `${Math.max(0, Math.min(1, frac)) * 100}%`;
ui.status.textContent = left; ui.statusR.textContent = right;
}
function showTimings(rows) {
ui.timings.innerHTML = rows.map(([k, v]) => `<dt>${k}</dt><dd>${v}</dd>`).join("");
}
function clearOutput() {
ui.placeholder.hidden = false; ui.aiNote.hidden = true; ui.thinker.hidden = true; ui.save.hidden = true; ui.prepared.hidden = true;
showTimings([]); last = null;
if (ui.save.href?.startsWith("blob:")) URL.revokeObjectURL(ui.save.href);
}
async function load() {
busy = true; ui.go.disabled = true;
window.__agate.state = "loading";
agate = new Agate({ base: BASE, override: MODELS_OVERRIDE, version, ep, variant: VARIANT, optLevel: params.get("opt") || "all" });
try {
const t0 = performance.now();
const st = await agate.load(({ phase, file, loaded, total }) => {
if (phase === "download") setProgress(loaded / total, `Downloading ${file}`, `${mb(loaded)} / ${mb(total)} MB`);
else setProgress(1, `Preparing ${file} on ${ep === "webgpu" ? "WebGPU" : "CPU"}`, `${mb(total)} MB`);
}, res());
const loadMs = performance.now() - t0;
$("dl-size").textContent = `${st.downloadMB.toFixed(0)} MB`;
ui.backend.textContent = `backend: ${ep === "webgpu" ? "WebGPU" : "WebAssembly (CPU)"}`;
const src = st.fromCacheMB > st.downloadMB * 0.99 ? "from browser cache" : "downloaded";
setProgress(0, "Ready", `Preview ${version} · ${st.downloadMB.toFixed(0)} MB ${src} · ${(loadMs / 1000).toFixed(1)} s`);
window.__agate = { state: "ready", version, load: { ...st, loadMs, ep } };
ui.go.textContent = "Generate"; ui.go.disabled = false;
} catch (e) {
console.error(e);
setProgress(0, "Load failed", String(e.message || e));
window.__agate = { state: "error", version, error: String(e.message || e) };
ui.go.textContent = "Retry load"; ui.go.disabled = false; agate = null;
}
busy = false;
}
async function pngWithMetadata() {
const blob = await new Promise((r) => ui.canvas.toBlob(r, "image/png"));
const bytes = addPngText(new Uint8Array(await blob.arrayBuffer()), agate.marks.info);
return new Blob([bytes], { type: "image/png" });
}
// The previews: frames go to thinkerView, which colours them in a Web Worker and drops frames while busy. The
// sampler never waits for any of it (agate.js sampleFast); with Show thinker off nothing is requested at all.
function previewFn() {
thinkerView.reset();
let n = 0, t0 = performance.now();
window.__agate.thinker = { submitted: 0 };
thinkerView.onDrawn = ({ step, steps, workerMs, dropped, drawn }) => {
ui.thinker.hidden = false;
ui.thinkerStep.textContent = `${step} / ${steps}`;
window.__agate.thinker = { submitted: n, drawn, dropped, lastWorkerMs: workerMs, lastStep: step };
};
return ({ step, steps, plan, x1, hw }) => {
if (thinkerView.submit(plan, x1, hw, { step, steps })) n++;
};
}
async function generate(opts = {}) {
busy = true; stop = false;
ui.go.textContent = "Stop"; ui.frame.classList.add("busy");
window.__agate.state = "generating";
const steps = opts.steps ?? Number(ui.steps.value);
try {
setProgress(0, "Encoding prompt", "");
const r = await agate.generate({
prompt: opts.prompt ?? ui.prompt.value.trim(), negative: opts.negative ?? ui.negative.value.trim(),
seed: Number(ui.seed.value) >>> 0, steps, cfg: opts.cfg ?? Number(ui.cfg.value), noise: opts.noise ?? null,
watermark: opts.watermark ?? true,
onStep: (i, n) => setProgress(i / n, `Step ${i} / ${n}`, ""),
onPreview: (opts.thinker ?? ui.showThinker.checked) ? previewFn() : null,
previewReady: () => !thinkerView.pending,
shouldStop: () => stop,
});
const ctx = ui.canvas.getContext("2d");
ui.canvas.width = r.width; ui.canvas.height = r.height;
ctx.putImageData(new ImageData(r.rgba, r.width, r.height), 0, 0);
ui.placeholder.hidden = true;
ui.aiNote.hidden = false;
const T = r.timings;
setProgress(1, "Done", `${(T.totalMs / 1000).toFixed(1)} s`);
showTimings([
["text encode", `${T.textMs.toFixed(0)} ms (bucket ${T.bucket})`],
["per step", `${T.stepMs.toFixed(0)} ms × ${steps} (first ${T.firstStepMs.toFixed(0)} ms)`],
["decode + mark", `${T.decodeMs.toFixed(0)} + ${T.markMs.toFixed(0)} ms`],
["total", `${(T.totalMs / 1000).toFixed(2)} s`],
]);
const p = r.prepared;
const changed = agate.mr && (p.text !== (opts.prompt ?? ui.prompt.value.trim()) || p.negative || p.countsNonzero);
ui.prepared.hidden = !changed;
if (changed) ui.prepared.textContent = `Model sees: ${p.text}${p.negative ? ` · avoid: ${p.negative}` : ""}${p.countsNonzero ? ` · count code on ${p.countsNonzero} token(s)` : ""}`;
if (ui.save.href?.startsWith("blob:")) URL.revokeObjectURL(ui.save.href);
ui.save.href = URL.createObjectURL(await pngWithMetadata());
ui.save.download = `agate-${version}-${r.width}px-seed${Number(ui.seed.value) >>> 0}.png`;
ui.save.hidden = false;
last = r;
window.__agate = { ...window.__agate, state: "done", timings: T, prepared: p, size: [r.width, r.height], watermarked: r.watermarked,
thinkerShown: !ui.thinker.hidden };
return r;
} catch (e) {
if (String(e.message) === "stopped") setProgress(0, "Stopped", "");
else { console.error(e); setProgress(0, "Error", String(e.message || e)); window.__agate.error = String(e.message || e); }
window.__agate.state = "done";
} finally {
busy = false; ui.go.textContent = "Generate"; ui.frame.classList.remove("busy");
}
}
ui.go.onclick = async () => {
if (busy && agate?.sessions) { stop = true; return; }
if (busy) return;
if (!agate) return load();
return generate();
};
ui.prompt.addEventListener("keydown", (e) => { if (e.key === "Enter" && (e.ctrlKey || e.metaKey) && agate && !busy) generate(); });
// Parity self-test: ?parity=1 runs the chosen version's fixture (export/parity_full*.py): same initial noise as the
// PyTorch package, watermark off (the fixture holds unmarked pixels); then checks the mark on the marked image.
async function parity() {
const dir = version === "001" ? "test/" : version === "002" ? "test/002/" : `test/003_${res()}/`;
const fx = await (await fetch(`${dir}parity.json`)).json();
const bin = async (f) => new Uint8Array(await (await fetch(`${dir}${f}`)).arrayBuffer());
const noise = new Float32Array((await bin("parity_noise.bin")).buffer);
const refZ = new Float32Array((await bin("parity_latent.bin")).buffer);
const refImg = await bin("parity_image.bin"); // HWC uint8
ui.prompt.value = fx.prompt; ui.steps.value = fx.steps; ui.cfg.value = fx.cfg; ui.negative.value = "";
const r = await generate({ prompt: fx.prompt, negative: "", steps: fx.steps, cfg: fx.cfg, noise, watermark: false });
let dz = 0, dzs = 0, di = 0, dis = 0, se = 0;
for (let k = 0; k < refZ.length; k++) { const d = Math.abs(r.latent[k] - refZ[k]); dz = Math.max(dz, d); dzs += d; }
for (let p = 0; p < refImg.length / 3; p++) for (let c = 0; c < 3; c++) {
const d = Math.abs(r.rgba[p * 4 + c] - refImg[p * 3 + c]); di = Math.max(di, d); dis += d; se += d * d;
}
const psnr = 10 * Math.log10(255 * 255 / Math.max(1e-9, se / refImg.length));
const prepOk = fx.prepared === undefined || (r.prepared.text === fx.prepared && r.prepared.negative === fx.negative);
// the same image marked, read back in the page
const { embedWatermark } = await import("./marking.js");
const marked = Uint8ClampedArray.from(r.rgba);
embedWatermark(marked, r.width, r.height, agate.marks.payload);
const rb = readWatermark(marked, r.width, r.height, agate.marks.payload);
const res_ = { version, resolution: res(), latentMaxAbs: dz, latentMeanAbs: dzs / refZ.length, imageMaxAbs: di, imageMeanAbs: dis / refImg.length,
psnr, preparedMatchesPython: prepOk, watermarkReadback: rb.text, watermarkBitAcc: rb.bitAccuracy, ep, timings: r.timings };
const el = $("parity"); el.hidden = false;
el.innerHTML = `<span class="label"><i class="sq red"></i>Parity vs PyTorch package · Preview ${version} · ${res()} px (${ep})</span>
<p>final latent max |Δ| ${dz.toFixed(4)} (mean ${res_.latentMeanAbs.toFixed(5)}) · image max |Δ| ${di}/255 (mean ${res_.imageMeanAbs.toFixed(3)}, PSNR ${psnr.toFixed(1)} dB)
· prompt pipeline ${prepOk ? "identical" : "DIFFERENT"} · watermark read back: ${rb.text} (${(rb.bitAccuracy * 100).toFixed(0)}% bits)</p>`;
window.__agate = { ...window.__agate, parity: res_, state: "parity-done" };
}
// Golden test of the prompt pipeline in this browser (tokenizers.js + prompt.js) against the Python package.
async function golden() {
const G = await (await fetch("test/prompt_golden.json")).json();
const tok = agate.tok, fails = [];
const enc = (s) => agate.tokenize(s).map(Number);
for (const c of G.cases) {
const [text, neg] = preparePrompt(c.prompt);
const ids = enc(text);
const cv = countVector(tok, text, ids, ids.length), counts = {};
cv.forEach((v, j) => { if (v) counts[String(j)] = v; });
const ok = text === c.prepared && neg === c.negative && JSON.stringify(ids) === JSON.stringify(c.ids)
&& JSON.stringify(enc(neg)) === JSON.stringify(c.neg_ids) && JSON.stringify(counts) === JSON.stringify(c.counts);
if (!ok) fails.push(c.prompt);
}
const res_ = { cases: G.cases.length, identical: G.cases.length - fails.length, fails };
const el = $("parity"); el.hidden = false;
el.innerHTML += `<span class="label"><i class="sq red"></i>Prompt pipeline golden test</span><p>${res_.identical} / ${res_.cases} prompts identical to the Python package (prompt, negative, token ids, count code)</p>`;
window.__agate = { ...window.__agate, golden: res_ };
}
// For the automated test: the PNG exactly as "Save PNG" gives it, base64
window.__agateSavedPng = async () => {
const b = new Uint8Array(await (await pngWithMetadata()).arrayBuffer());
let s = ""; for (let i = 0; i < b.length; i += 0x8000) s += String.fromCharCode(...b.subarray(i, i + 0x8000));
return btoa(s);
};
window.agateClearCache = clearCache;
(async () => {
renderVersion();
if (!FORCE_EP) {
const s = await webgpuStatus();
if (!s.ok) {
gpuOk = false; ep = "wasm";
$("nogpu").hidden = false; $("nogpu-why").textContent = s.why;
ui.go.disabled = true; ui.go.textContent = "WebGPU required";
window.__agate = { state: "no-webgpu", why: s.why, version };
$("use-wasm").onclick = () => { $("nogpu").hidden = true; gpuOk = true; ui.go.textContent = "Load model (CPU)"; ui.go.disabled = false; };
return;
}
ui.backend.textContent = `backend: WebGPU${s.adapter ? " · " + s.adapter : ""}`;
}
ui.go.disabled = false;
ui.go.textContent = "Load model";
if (params.has("autoload") || PARITY || GOLDEN || await hasCachedModel(version)) { // second visit: load straight from cache
await load();
if (GOLDEN && agate && version === "003") await golden();
if (PARITY && agate) await parity();
if (GOLDEN && !PARITY) window.__agate.state = "golden-done";
}
})();