Spaces:
Running
Running
File size: 20,798 Bytes
2dc8d66 7a7d8dc 2dc8d66 7a7d8dc 2dc8d66 146ba97 2dc8d66 7a7d8dc 2dc8d66 7a7d8dc 2dc8d66 7a7d8dc 2dc8d66 7a7d8dc 2dc8d66 7a7d8dc b0dc8c9 2dc8d66 b0dc8c9 7a7d8dc 2dc8d66 7a7d8dc 2dc8d66 7a7d8dc 2dc8d66 146ba97 2dc8d66 7a7d8dc 2dc8d66 7a7d8dc 2dc8d66 146ba97 2dc8d66 7a7d8dc 2dc8d66 7a7d8dc 2dc8d66 7a7d8dc 2dc8d66 7a7d8dc 2dc8d66 7a7d8dc 2dc8d66 7a7d8dc 2dc8d66 7a7d8dc 2dc8d66 146ba97 2dc8d66 7a7d8dc 2dc8d66 146ba97 2dc8d66 146ba97 2dc8d66 7a7d8dc 2dc8d66 7a7d8dc 2dc8d66 146ba97 2dc8d66 146ba97 2dc8d66 7a7d8dc 2dc8d66 7a7d8dc 2dc8d66 7a7d8dc 2dc8d66 7a7d8dc 2dc8d66 7a7d8dc 146ba97 7a7d8dc 2dc8d66 7a7d8dc 2dc8d66 7a7d8dc 2dc8d66 7a7d8dc 2dc8d66 7a7d8dc 146ba97 7a7d8dc 2dc8d66 7a7d8dc 2dc8d66 7a7d8dc | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 280 281 282 283 284 285 286 287 288 289 290 291 292 293 294 295 296 297 298 299 300 301 302 303 304 305 306 307 308 309 310 311 312 313 314 315 316 317 318 319 320 321 322 323 324 325 326 327 328 329 330 331 332 333 334 335 336 337 338 339 340 341 342 343 344 345 346 347 348 349 350 351 352 353 354 355 356 357 358 359 360 361 362 363 364 365 366 367 368 369 370 371 372 373 374 375 376 377 378 379 380 381 382 383 384 385 386 387 388 389 | // 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 };
}
}
|