// Agate in the browser: tokenizer (tokenizers.js) + ONNX graphs run by onnxruntime-web (WebGPU, WASM // fallback) + an Euler flow sampler with CFG. Mirrors agate/pipeline.py (AgatePipeline) of // Logolabs/agate-preview-001 / -002 (fcdm_t2, 256 px) and Logolabs/agate-preview-003 (fcdm_t2mr: 512 or 256 px, // SD3 timestep shift, prompt pipeline + count code, see prompt.js). Every image is marked (marking.js). import * as ort from "https://cdn.jsdelivr.net/npm/onnxruntime-web@1.30.0/dist/ort.webgpu.bundle.min.mjs"; import { Tokenizer } from "https://cdn.jsdelivr.net/npm/@huggingface/tokenizers@0.2.0/dist/tokenizers.mjs"; import { prepare, countVector } from "./prompt.js"; import { embedWatermark, marks } from "./marking.js"; export { ort }; // The releases this page runs. remote: where the files live -- each release's model repo, under webgpu/ (the // Space repo has a 1 GB storage limit; the graphs there also output the thinker plan). The Space's old models/ // (001 without the plan output) is used only through ?models=. dir: the layout under a local // models base (?models= or window.AGATE_MODEL_BASE given explicitly, e.g. for local testing). cache: Cache // Storage name (001 keeps the name it always had, so returning visitors do not download it again). const HUB = (v) => `https://huggingface.co/Logolabs/agate-preview-${v}/resolve/main/webgpu/v2/`; // v2: fast fp16 step graphs export const VERSIONS = { "003": { remote: HUB("003"), dir: "003/", cache: "agate-preview-003-web", res: [512, 256], note: "Multi-resolution model, 512 px (or 256). Quoted text is spelled out, counts are coded, “no X” becomes a negative prompt." }, "002": { remote: HUB("002"), dir: "002/", cache: "agate-preview-002-web", res: [256], note: "Same network as 001, trained 38,710 steps longer before the same anneal. 256 px, plain prompts." }, "001": { remote: HUB("001"), dir: "", cache: "agate-preview-001-web", res: [256], note: "The first preview (2026-09-25). 256 px, plain prompts." }, }; const countUrl = (v) => `https://huggingface.co/Logolabs/agate-preview-${v}/resolve/main/config.json`; const FILES_T2 = ["tokenizer.json", "tokenizer_config.json", "text_encoder.onnx", "generator.onnx", "taesd_decoder.onnx"]; export async function webgpuStatus() { if (!("gpu" in navigator)) return { ok: false, why: "This browser does not expose WebGPU (navigator.gpu is missing)." }; try { const adapter = await navigator.gpu.requestAdapter({ powerPreference: "high-performance" }); if (!adapter) return { ok: false, why: "WebGPU is present but no GPU adapter is available." }; let name = ""; try { const info = adapter.info || (await adapter.requestAdapterInfo?.()); name = [info?.vendor, info?.architecture, info?.description].filter(Boolean).join(" "); } catch { /* optional */ } return { ok: true, adapter: name, maxBuffer: adapter.limits?.maxBufferSize }; } catch (e) { return { ok: false, why: `WebGPU adapter request failed: ${e.message || e}` }; } } // ---- downloads with progress + Cache Storage --------------------------------------------------- async function openCache(name) { try { return await caches.open(name); } catch { return null; } // file://, private mode, ... } async function fetchCached(cacheName, url, key, expectBytes, onBytes) { const cache = await openCache(cacheName); if (cache) { const hit = await cache.match(key); if (hit) { const buf = new Uint8Array(await hit.arrayBuffer()); onBytes(buf.byteLength, true); return { buf, cached: true }; } } const res = await fetch(url); if (!res.ok) throw new Error(`${url}: HTTP ${res.status}`); // manifest size first: a compressed response would report the encoded length const total = expectBytes || Number(res.headers.get("content-length")) || 0; const reader = res.body.getReader(); let buf = new Uint8Array(total || 1 << 20), n = 0; for (;;) { const { done, value } = await reader.read(); if (done) break; if (n + value.byteLength > buf.byteLength) { const b2 = new Uint8Array(Math.max(buf.byteLength * 2, n + value.byteLength)); b2.set(buf.subarray(0, n)); buf = b2; } buf.set(value, n); n += value.byteLength; onBytes(value.byteLength, false); } buf = n === buf.byteLength ? buf : buf.slice(0, n); if (cache) { try { await cache.put(key, new Response(buf, { headers: { "content-type": "application/octet-stream", "content-length": String(n) } })); } catch (e) { console.warn("cache put failed (quota?)", e); } } return { buf, cached: false }; } export async function hasCachedModel(version = "001") { try { const c = await caches.open(VERSIONS[version].cache); return (await c.keys()).length >= 5; } catch { return false; } } export async function clearCache() { let ok = true; for (const v of Object.values(VERSIONS)) { try { ok = (await caches.delete(v.cache)) && ok; } catch { ok = false; } } return ok; } // The Hub counts a model download only on a request for the model repo's root config.json; the ONNX files // come from this Space, so every model load reads the loaded release's config.json once -- the same per-load // count a Python from_pretrained() produces, cached weights or not. function countLoad(version) { fetch(countUrl(version), { cache: "no-store" }).catch(() => {}); } // ---- the model --------------------------------------------------------------------------------- export class Agate { // base: the page's models folder (001); override: an explicit models base for every version (local layout) constructor({ base = "./models/", override = null, version = "003", ep = "webgpu", variant = null, optLevel = "all" } = {}) { this.optLevel = optLevel; const slash = (u) => (u.endsWith("/") ? u : u + "/"); this.version = version; this.spec = VERSIONS[version]; this.base = override ? slash(override) + this.spec.dir : (this.spec.remote || slash(base) + this.spec.dir); this.ep = ep; this.variant = variant; // e.g. "full16": fp16-compute generator (experimental, 001 only) this.marks = marks(version); } get mr() { return this.manifest?.arch === "fcdm_t2mr"; } // the fast build (2026-09-29): one Euler step per run, fp16 compute, everything stays on the GPU between steps get fast() { return !!this.manifest?.step; } fileKey(f) { return new Request(`${this.base}${f}?sha256=${this.manifest.files[f].sha256}`); } async getFile(f, onBytes = () => {}) { const meta = this.manifest.files[f]; return (await fetchCached(this.spec.cache, this.base + f, this.fileKey(f), meta.bytes, onBytes)).buf; } // onProgress({loaded, total, file, phase}); resolution: 003 only (512 default, or 256) async load(onProgress = () => {}, resolution = null) { const manifest = await (await fetch(this.base + "manifest.json", { cache: "no-cache" })).json(); this.manifest = manifest; const swap = (this.variant && manifest.variants?.[this.variant]) || {}; const real = (f) => swap[f] || f; this.resolution = this.mr ? Number(resolution || manifest.resolution) : 256; // 003: download everything (the two graphs are ~1 MB each and share one weights file), so a resolution // switch needs no network const need = this.mr || manifest.step ? Object.keys(manifest.files) : FILES_T2; const total = need.reduce((s, f) => s + manifest.files[real(f)].bytes, 0); let loaded = 0, fromCache = 0; const bufs = {}; const t0 = performance.now(); for (const f of need) { const meta = manifest.files[real(f)]; const { buf } = await fetchCached(this.spec.cache, this.base + real(f), this.fileKey(real(f)), meta.bytes, (b, c) => { loaded += b; if (c) fromCache += b; onProgress({ phase: "download", file: f, loaded, total }); }); bufs[f] = buf; } this.stats = { downloadMB: total / 1e6, fromCacheMB: fromCache / 1e6, downloadMs: performance.now() - t0 }; countLoad(this.version); // drop cached files of older builds of this release (keys carry the sha256) try { const cache = await openCache(this.spec.cache), keep = new Set(need.map((f) => new URL(this.fileKey(real(f)).url, location.href).href)); if (cache) for (const req of await cache.keys()) if (!keep.has(req.url)) await cache.delete(req); } catch (e) { console.warn("cache cleanup", e); } const dec = new TextDecoder(); this.tok = new Tokenizer(JSON.parse(dec.decode(bufs["tokenizer.json"])), JSON.parse(dec.decode(bufs["tokenizer_config.json"]))); if (this.ep === "wasm") { ort.env.wasm.numThreads = self.crossOriginIsolated ? Math.min(8, navigator.hardwareConcurrency || 4) : 1; } const t1 = performance.now(); this.sessions = {}; for (const [name, f] of [["text", "text_encoder.onnx"], ["vae", "taesd_decoder.onnx"]]) { onProgress({ phase: "init", file: f, loaded: total, total }); const ts = performance.now(); this.sessions[name] = await ort.InferenceSession.create(bufs[f], this.sessionOptions()); this.stats[`${name}SessionMs`] = performance.now() - ts; bufs[f] = null; } onProgress({ phase: "init", file: "generator", loaded: total, total }); const ts = performance.now(); if (this.fast) { this.weights = bufs[manifest.step.external_data]; // kept for a resolution switch (no re-read) await this.createGenerator(this.resolution, bufs); } else if (this.mr) { this.weights = bufs["generator.weights"]; // kept for a resolution switch (no re-read) await this.createGenerator(this.resolution, bufs); } else { this.sessions.gen = await ort.InferenceSession.create(bufs["generator.onnx"], this.sessionOptions()); } this.stats.genSessionMs = performance.now() - ts; this.stats.sessionMs = performance.now() - t1; return this.stats; } sessionOptions(extra = {}) { return { executionProviders: [this.ep], graphOptimizationLevel: this.optLevel, ...extra }; } // 003: the generator graph for one resolution; the weights file is shared by both graphs. async createGenerator(res, bufs = {}) { let graphName, ext, extra = {}; if (this.fast) { graphName = this.manifest.step.graphs[String(res)]; ext = this.manifest.step.external_data; extra = { preferredOutputLocation: "gpu-buffer" }; // z_next / x1 / plan stay on the GPU } else { const r = this.manifest.resolutions[String(res)]; graphName = r?.graph; ext = r?.external_data; } if (!graphName) throw new Error(`resolution ${res} not available`); const graph = bufs[graphName] || await this.getFile(graphName); const weights = this.weights || await this.getFile(ext); if (this.sessions.gen) { try { await this.sessions.gen.release(); } catch { /* */ } this.sessions.gen = null; } this.sessions.gen = await ort.InferenceSession.create(graph, this.sessionOptions({ externalData: [{ path: ext, data: weights }], ...extra })); this.resolution = Number(res); } async setResolution(res) { if (!this.mr || Number(res) === this.resolution) return 0; const t0 = performance.now(); await this.createGenerator(res); return performance.now() - t0; } async release() { for (const s of Object.values(this.sessions || {})) { try { await s?.release(); } catch { /* */ } } this.sessions = null; this.weights = null; } tokenize(text) { const max = this.manifest.text_max_len; let ids = this.tok.encode(text).ids; // [CLS] ... [SEP], as the HF tokenizer if (ids.length > max) ids = ids.slice(0, max - 1).concat(ids[ids.length - 1]); // truncation=True return ids; } async encode(text) { const ids = this.tokenize(text); const n = ids.length; const feeds = { input_ids: new ort.Tensor("int64", BigInt64Array.from(ids, BigInt), [1, n]), attention_mask: new ort.Tensor("int64", new BigInt64Array(n).fill(1n), [1, n]), }; const out = await this.sessions.text.run(feeds); const h = out.last_hidden_state; const data = h.data.slice(); // (1, n, 512) fp32 h.dispose?.(); return { data, n, ids }; } // The prompt as the model sees it: 003 runs its prompt pipeline (prompt.js), 001/002 use the prompt as typed. prepare(prompt, negative = "") { const pp = this.manifest.prompt_pipeline; if (!pp) return { text: prompt, negative }; const [text, neg] = prepare(prompt, negative, pp.normalize, pp.spell); return { text, negative: neg }; } static gaussianNoise(seed, count) { // mulberry32 + Box-Muller. Seeds do NOT reproduce PyTorch's torch.randn stream. let a = seed >>> 0; const rnd = () => { a |= 0; a = (a + 0x6D2B79F5) | 0; let t = Math.imul(a ^ (a >>> 15), 1 | a); t = (t + Math.imul(t ^ (t >>> 7), 61 | t)) ^ t; return ((t ^ (t >>> 14)) >>> 0) / 4294967296; }; const out = new Float32Array(count); for (let i = 0; i < count; i += 2) { const u1 = Math.max(rnd(), 1e-12), u2 = rnd(); const r = Math.sqrt(-2 * Math.log(u1)); out[i] = r * Math.cos(2 * Math.PI * u2); if (i + 1 < count) out[i + 1] = r * Math.sin(2 * Math.PI * u2); } return out; } // SD3's resolution shift in Agate's convention (t = 0 noise), as agate/pipeline.py shift_t static shiftT(t, shift) { if (shift === 1) return t; const s = 1 - t; return 1 - shift * s / (1 + (shift - 1) * s); } // -> { rgba (watermarked unless watermark=false), width, height, latent, timings, prepared } get hasPlan() { return !!this.sessions?.gen?.outputNames?.includes("plan"); } // GPU-resident sampler for the step graphs: z_next of one run is the z input of the next, ctx / mask / counts // are uploaded once, and nothing is read back until the end. The preview (plan + x1) is requested only when // onPreview is set, previewReady() says the viewer is free and no earlier readback is still in flight; its // download is started but never awaited by the loop (frames are dropped instead). async sampleFast({ z, grid, cfg, hw, C, L, D, ctx, mask, cnt, onStep, onPreview, previewReady, shouldStop, stepMs }) { const dev = ort.env.webgpu?.device; const bufs = []; const gpuT = (data, dims) => { if (!dev) return new ort.Tensor("float32", data, dims); const buf = dev.createBuffer({ size: Math.ceil(data.byteLength / 16) * 16, usage: GPUBufferUsage.STORAGE | GPUBufferUsage.COPY_SRC | GPUBufferUsage.COPY_DST }); dev.queue.writeBuffer(buf, 0, data); bufs.push(buf); return ort.Tensor.fromGpuBuffer(buf, { dataType: "float32", dims }); }; const fixed = { ctx: gpuT(ctx, [2, L, D]), mask: gpuT(mask, [2, L]) }; if (this.sessions.gen.inputNames.includes("counts")) fixed.counts = gpuT(cnt || new Float32Array(2 * L), [2, L]); const steps = grid.length - 1; let zT = new ort.Tensor("float32", z, [1, C, hw, hw]), inflight = false; const one = (v) => new ort.Tensor("float32", new Float32Array([v]), [1]); const cfgT = one(cfg); try { for (let i = 0; i < steps; i++) { if (shouldStop()) throw new Error("stopped"); const ts = performance.now(); const want = !!onPreview && this.hasPlan && !inflight && previewReady(); const out = await this.sessions.gen.run({ z: zT, t: one(grid[i]), dt: one(grid[i + 1] - grid[i]), cfg: cfgT, ...fixed }, want ? ["z_next", "x1", "plan"] : ["z_next"]); if (i > 0) zT.dispose?.(); zT = out.z_next; if (want) { inflight = true; const step = i + 1; Promise.all([out.plan.getData(true), out.x1.getData(true)]) .then(([plan, x1]) => onPreview({ step, steps, plan, x1, hw })) .catch((e) => console.warn("preview", e)) .finally(() => { inflight = false; }); } stepMs.push(performance.now() - ts); onStep(i + 1, steps, null); if ((i & 3) === 3) await new Promise((r) => setTimeout(r, 0)); // let the page paint now and then } const zf = await zT.getData(true); // the only synchronising readback return Float32Array.from(zf); } finally { for (const b of bufs) { try { b.destroy(); } catch { /* */ } } } } // onPreview({ step, steps, plan, x1, hw }) after every step, when given and the graph outputs the plan: plan = // the thinker output for the conditional branch (Float32Array 640 x 16 x 16), x1 = z + (1 - t) v (guided). async generate({ prompt, negative = "", seed = 0, steps = 50, cfg = 3.0, noise = null, watermark = true, onStep = () => {}, onPreview = null, previewReady = () => true, shouldStop = () => false }) { const C = this.manifest.latent_ch, D = this.manifest.ctx_dim; const rc = this.mr ? this.manifest.resolutions[String(this.resolution)] : { latent_hw: this.manifest.latent_hw, shift: 1 }; const hw = rc.latent_hw, shift = Number(rc.shift || 1); const per = C * hw * hw; const T = { }; let t0 = performance.now(); const prep = this.prepare(prompt, negative); const c = await this.encode(prep.text), u = await this.encode(prep.negative); T.textMs = performance.now() - t0; const BUCKETS = this.manifest.buckets; const need = Math.max(c.n, u.n); const L = BUCKETS.find((b) => b >= need) ?? BUCKETS[BUCKETS.length - 1]; const ctx = new Float32Array(2 * L * D), mask = new Float32Array(2 * L); ctx.set(c.data, 0); ctx.set(u.data, L * D); // zero padding, as F.pad in the pipeline mask.fill(1, 0, c.n); mask.fill(1, L, L + u.n); const feeds = { ctx: new ort.Tensor("float32", ctx, [2, L, D]), mask: new ort.Tensor("float32", mask, [2, L]) }; let countsNonzero = 0; if (this.mr) { // count code: conditional half only const cnt = new Float32Array(2 * L); if (this.manifest.prompt_pipeline?.count_code) cnt.set(countVector(this.tok, prep.text, c.ids, Math.min(L, c.n)), 0); countsNonzero = cnt.reduce((s, v) => s + (v > 0), 0); feeds.counts = new ort.Tensor("float32", cnt, [2, L]); } let z = noise ? Float32Array.from(noise) : Agate.gaussianNoise(seed, per); const grid = Array.from({ length: steps + 1 }, (_, i) => Agate.shiftT(i / steps, shift)); t0 = performance.now(); const stepMs = []; if (this.fast) { z = await this.sampleFast({ z, grid, cfg, hw, C, L, D, ctx, mask, cnt: feeds.counts?.data, onStep, onPreview, previewReady, shouldStop, stepMs }); } else { const zz = new Float32Array(2 * per), tt = new Float32Array(2); for (let i = 0; i < steps; i++) { if (shouldStop()) throw new Error("stopped"); const ts = performance.now(); zz.set(z, 0); zz.set(z, per); tt[0] = tt[1] = grid[i]; const want = onPreview && this.hasPlan ? ["v", "plan"] : ["v"]; const out = await this.sessions.gen.run({ z: new ort.Tensor("float32", zz, [2, C, hw, hw]), t: new ort.Tensor("float32", tt, [2]), ...feeds }, want); const v = out.v.data; const dt = grid[i + 1] - grid[i]; const zn = new Float32Array(per); let x1 = null; if (out.plan) x1 = new Float32Array(per); for (let k = 0; k < per; k++) { const vc = v[k], vu = v[per + k], g = vu + cfg * (vc - vu); zn[k] = z[k] + dt * g; if (x1) x1[k] = z[k] + (1 - grid[i]) * g; } out.v.dispose?.(); if (out.plan) { const plan = out.plan.data.slice(); out.plan.dispose?.(); onPreview({ step: i + 1, steps, plan, x1, hw }); } z = zn; stepMs.push(performance.now() - ts); onStep(i + 1, steps, z); await new Promise((r) => setTimeout(r, 0)); // let the progress bar paint } } T.samplerMs = performance.now() - t0; T.firstStepMs = stepMs[0]; T.stepMs = stepMs.length > 1 ? (T.samplerMs - stepMs[0]) / (stepMs.length - 1) : stepMs[0]; t0 = performance.now(); const img = await this.sessions.vae.run({ latent: new ort.Tensor("float32", z, [1, C, hw, hw]) }); const x = img.image.data, H = img.image.dims[2], W = img.image.dims[3]; const rgba = new Uint8ClampedArray(H * W * 4); for (let p = 0; p < H * W; p++) { for (let ch = 0; ch < 3; ch++) { const val = Math.min(1, Math.max(-1, x[ch * H * W + p])); rgba[p * 4 + ch] = Math.round((val + 1) * 127.5); } rgba[p * 4 + 3] = 255; } T.decodeMs = performance.now() - t0; t0 = performance.now(); if (watermark) embedWatermark(rgba, W, H, this.marks.payload); T.markMs = performance.now() - t0; T.totalMs = T.textMs + T.samplerMs + T.decodeMs + T.markMs; T.bucket = L; return { rgba, width: W, height: H, latent: z, timings: T, prepared: { ...prep, countsNonzero }, watermarked: watermark }; } }