Spaces:
Running
Running
Stefatorus
Claude Opus 5.5
Faster WebGPU sampling: fp16 compute, GPU-resident steps, non-blocking thinker preview
146ba97 Download js/agate.js from Logolabs/agate-webgpu: direct link, hf CLI and curl.
- Browser
- Download file 20.8 kB
-
https://huggingface.co/spaces/Logolabs/agate-webgpu/resolve/main/js/agate.js
- Command line
-
hf download hf://spaces/Logolabs/agate-webgpu/js/agate.js
-
curl -L -o agate.js https://huggingface.co/spaces/Logolabs/agate-webgpu/resolve/main/js/agate.js
20.8 kB
| // 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=<url> 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 }; | |
| } | |
| } | |