ysharma HF Staff commited on
Commit
3c4c02e
·
verified ·
1 Parent(s): de28bb9

Agate 4-step WebGPU: in-browser runtime with live mode

Browse files
Files changed (7) hide show
  1. README.md +35 -5
  2. css/style.css +79 -0
  3. index.html +71 -17
  4. js/agate.js +390 -0
  5. js/app.js +169 -0
  6. js/marking.js +188 -0
  7. js/prompt.js +249 -0
README.md CHANGED
@@ -1,10 +1,40 @@
1
  ---
2
- title: Agate 4step Webgpu
3
- emoji: 🐨
4
- colorFrom: blue
5
- colorTo: yellow
6
  sdk: static
 
7
  pinned: false
 
 
 
 
 
 
 
 
 
8
  ---
9
 
10
- Check out the configuration reference at https://huggingface.co/docs/hub/spaces-config-reference
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
  ---
2
+ title: Agate 4-step WebGPU
3
+ emoji: ⚡
4
+ colorFrom: red
5
+ colorTo: gray
6
  sdk: static
7
+ app_file: index.html
8
  pinned: false
9
+ license: mit
10
+ short_description: 4-step Agate text-to-image in your browser, live as you type
11
+ custom_headers:
12
+ cross-origin-embedder-policy: require-corp
13
+ cross-origin-opener-policy: same-origin
14
+ cross-origin-resource-policy: cross-origin
15
+ models:
16
+ - ysharma/agate-002-4step
17
+ - Logolabs/agate-preview-002
18
  ---
19
 
20
+ # Agate 4-step, in your browser
21
+
22
+ [ysharma/agate-002-4step](https://huggingface.co/ysharma/agate-002-4step) is an **unofficial** 4-step distillation of
23
+ [Agate Preview 002](https://huggingface.co/Logolabs/agate-preview-002) by LogoLabs. This page runs it entirely on your
24
+ GPU with WebGPU (onnxruntime-web). Prompts never leave your machine.
25
+
26
+ - **4 network passes per image** instead of Agate's 100: guidance is baked into the weights, and the sampler is 4 uniform Euler steps.
27
+ - **Live mode** redraws the image on every keystroke. Keystrokes typed while an image is being made are coalesced, so only the newest prompt gets drawn next.
28
+ - The model files (about 615 MB) download once from the model repo's `webgpu/` folder and are then kept in the browser's Cache Storage.
29
+
30
+ The runtime (tokenizer, text encoder, sampler, TAESD decoder, watermark) is adapted from LogoLabs'
31
+ [agate-webgpu](https://huggingface.co/spaces/Logolabs/agate-webgpu) (MIT). The one change to the sampler: with `cfg = 1`
32
+ the generator runs a batch of 1, not a conditional/unconditional pair.
33
+
34
+ **Limitations.** 256 × 256 only; no negative prompts; Agate's weaknesses (exact text, counts above three, negation) remain.
35
+ No safety filter runs in this page, and the base model can draw people partly unclothed without being asked.
36
+ Every image carries Agate's invisible watermark (`AGATE002`), and saved PNGs include AI-generated provenance fields.
37
+
38
+ Query parameters: `?ep=wasm` forces the CPU backend, `?autoload=1` loads the model on open.
39
+
40
+ Not affiliated with LogoLabs. Licence: MIT.
css/style.css ADDED
@@ -0,0 +1,79 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ :root {
2
+ --bg: #f7f6f3;
3
+ --surface: #ffffff;
4
+ --text: #1d1c1a;
5
+ --muted: #6b6862;
6
+ --line: #e2dfd8;
7
+ --accent: #b4462f;
8
+ --accent-text: #ffffff;
9
+ --frame: #ecE9e2;
10
+ }
11
+ @media (prefers-color-scheme: dark) {
12
+ :root:not([data-theme="light"]) {
13
+ --bg: #151412;
14
+ --surface: #1f1d1a;
15
+ --text: #efece6;
16
+ --muted: #a39f97;
17
+ --line: #34312c;
18
+ --accent: #e0765d;
19
+ --accent-text: #1a0f0c;
20
+ --frame: #26231f;
21
+ }
22
+ }
23
+ :root[data-theme="dark"] {
24
+ --bg: #151412; --surface: #1f1d1a; --text: #efece6; --muted: #a39f97;
25
+ --line: #34312c; --accent: #e0765d; --accent-text: #1a0f0c; --frame: #26231f;
26
+ }
27
+ * { box-sizing: border-box; }
28
+ html, body { margin: 0; background: var(--bg); color: var(--text); }
29
+ body { font: 15px/1.55 Inter, system-ui, sans-serif; }
30
+ a { color: var(--accent); }
31
+ .page { max-width: 880px; margin: 0 auto; padding: 32px 16px 64px; }
32
+ h1 { font-size: 28px; line-height: 1.2; margin: 0 0 8px; font-weight: 700; letter-spacing: -0.01em; }
33
+ h2 { font-size: 17px; margin: 0 0 8px; }
34
+ .tag { font: 500 12px/1 "JetBrains Mono", monospace; vertical-align: middle; padding: 4px 8px; border-radius: 999px;
35
+ border: 1px solid var(--line); color: var(--muted); }
36
+ .lede { color: var(--muted); margin: 0 0 24px; max-width: 64ch; }
37
+ .panel, .result, .about, .notice { background: var(--surface); border: 1px solid var(--line); border-radius: 12px; padding: 16px; margin-bottom: 16px; }
38
+ .notice { border-color: var(--accent); }
39
+ .label { display: block; font-weight: 600; font-size: 13px; margin: 12px 0 6px; }
40
+ textarea { width: 100%; resize: vertical; font: inherit; color: var(--text); background: var(--bg);
41
+ border: 1px solid var(--line); border-radius: 8px; padding: 10px 12px; }
42
+ textarea:focus, input:focus { outline: 2px solid var(--accent); outline-offset: 1px; }
43
+ button { font: 600 14px/1 Inter, system-ui, sans-serif; border-radius: 8px; padding: 10px 14px; cursor: pointer; }
44
+ button:disabled { opacity: 0.5; cursor: default; }
45
+ .primary { background: var(--accent); color: var(--accent-text); border: 1px solid var(--accent); }
46
+ .ghost { background: transparent; color: var(--text); border: 1px solid var(--line); }
47
+ .loadrow { margin-bottom: 10px; }
48
+ .loadrow .primary { width: 100%; }
49
+ .bar { height: 4px; background: var(--line); border-radius: 2px; overflow: hidden; }
50
+ #bar-fill { height: 100%; width: 0; background: var(--accent); transition: width 0.15s; }
51
+ .status { display: flex; justify-content: space-between; gap: 12px; font: 12px/1.4 "JetBrains Mono", monospace; color: var(--muted); margin: 6px 0 0; }
52
+ .chips { display: flex; flex-wrap: wrap; gap: 6px; margin-top: 8px; }
53
+ .chip { font: 12px/1.3 Inter, system-ui, sans-serif; padding: 6px 10px; border-radius: 999px; border: 1px solid var(--line);
54
+ background: transparent; color: var(--muted); }
55
+ .chip:hover { color: var(--text); border-color: var(--muted); }
56
+ .controls { display: flex; flex-wrap: wrap; align-items: center; gap: 10px; margin-top: 14px; }
57
+ .controls .primary { margin-left: auto; }
58
+ .switch { display: inline-flex; align-items: center; gap: 8px; font-weight: 500; }
59
+ .switch input { width: 16px; height: 16px; accent-color: var(--accent); }
60
+ .seed { display: inline-flex; align-items: center; gap: 6px; color: var(--muted); font-size: 13px; }
61
+ .seed input { width: 120px; font: 13px "JetBrains Mono", monospace; color: var(--text); background: var(--bg);
62
+ border: 1px solid var(--line); border-radius: 8px; padding: 8px; }
63
+ .result { display: grid; grid-template-columns: minmax(0, 512px) minmax(0, 1fr); gap: 16px; align-items: start; }
64
+ .frame { position: relative; aspect-ratio: 1; background: var(--frame); border-radius: 8px; overflow: hidden; }
65
+ .frame canvas { width: 100%; height: 100%; display: block; image-rendering: auto; }
66
+ .placeholder { position: absolute; inset: 0; display: grid; place-items: center; color: var(--muted); font-size: 14px; }
67
+ .frame.busy canvas { opacity: 0.85; }
68
+ .timings { display: grid; grid-template-columns: auto auto; gap: 4px 12px; margin: 0 0 12px; font: 12px/1.5 "JetBrains Mono", monospace; }
69
+ .timings dt { color: var(--muted); }
70
+ .timings dd { margin: 0; text-align: right; font-variant-numeric: tabular-nums; }
71
+ .timings .big { font-size: 22px; color: var(--text); font-weight: 500; }
72
+ .ai { font-size: 12px; color: var(--muted); }
73
+ .about p { color: var(--muted); max-width: 72ch; }
74
+ .fine { font-size: 12px; }
75
+ @media (max-width: 640px) {
76
+ .result { grid-template-columns: 1fr; }
77
+ .controls .primary { margin-left: 0; width: 100%; }
78
+ h1 { font-size: 24px; }
79
+ }
index.html CHANGED
@@ -1,19 +1,73 @@
1
  <!doctype html>
2
- <html>
3
- <head>
4
- <meta charset="utf-8" />
5
- <meta name="viewport" content="width=device-width" />
6
- <title>My static Space</title>
7
- <link rel="stylesheet" href="style.css" />
8
- </head>
9
- <body>
10
- <div class="card">
11
- <h1>Welcome to your static Space!</h1>
12
- <p>You can modify this app directly by editing <i>index.html</i> in the Files and versions tab.</p>
13
- <p>
14
- Also don't forget to check the
15
- <a href="https://huggingface.co/docs/hub/spaces" target="_blank">Spaces documentation</a>.
16
- </p>
17
- </div>
18
- </body>
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
19
  </html>
 
1
  <!doctype html>
2
+ <html lang="en">
3
+ <head>
4
+ <meta charset="utf-8">
5
+ <meta name="viewport" content="width=device-width, initial-scale=1">
6
+ <title>Agate 4-step WebGPU</title>
7
+ <meta name="description" content="An unofficial 4-step distillation of Agate Preview 002 (LogoLabs) running entirely in your browser on WebGPU. Redraws as you type.">
8
+ <link rel="preconnect" href="https://fonts.googleapis.com">
9
+ <link rel="preconnect" href="https://fonts.gstatic.com" crossorigin>
10
+ <link crossorigin="anonymous" href="https://fonts.googleapis.com/css2?family=Inter:wght@400;500;600;700&family=JetBrains+Mono:wght@400;500&display=swap" rel="stylesheet">
11
+ <link rel="stylesheet" href="css/style.css">
12
+ </head>
13
+ <body>
14
+ <main class="page">
15
+ <header>
16
+ <h1>Agate 4-step <span class="tag">WebGPU</span></h1>
17
+ <p class="lede">A 0.19B text-to-image model distilled to <b>4 steps with no guidance pass</b>, running on your own GPU.
18
+ Nothing leaves your machine. Turn on <b>Live</b> and the image redraws as you type.</p>
19
+ </header>
20
+
21
+ <div id="nogpu" class="notice" hidden>
22
+ <p id="nogpu-why"></p>
23
+ <p>This page needs WebGPU (recent Chrome or Edge). You can still run it on the CPU, which takes a few seconds per image.</p>
24
+ <button id="use-wasm" class="ghost">Use CPU instead</button>
25
+ </div>
26
+
27
+ <section class="panel">
28
+ <div class="loadrow" id="loadrow">
29
+ <button id="load" class="primary">Load model (about 615 MB, cached after the first visit)</button>
30
+ </div>
31
+ <div class="bar" aria-hidden="true"><div id="bar-fill"></div></div>
32
+ <p class="status"><span id="status">Not loaded</span><span id="status-r"></span></p>
33
+
34
+ <label for="prompt" class="label">Prompt</label>
35
+ <textarea id="prompt" rows="3" spellcheck="false" placeholder="a red cube on top of a blue sphere">a lighthouse on a rocky coast under a stormy sky</textarea>
36
+ <div class="chips" id="examples"></div>
37
+
38
+ <div class="controls">
39
+ <label class="switch"><input type="checkbox" id="live" checked> <span>Live (redraw on every keystroke)</span></label>
40
+ <label class="seed">Seed <input id="seed" type="number" value="0" min="0" max="4294967295"></label>
41
+ <button id="dice" class="ghost" title="New seed">New seed</button>
42
+ <button id="go" class="primary" disabled>Generate</button>
43
+ </div>
44
+ </section>
45
+
46
+ <section class="result">
47
+ <div class="frame" id="frame">
48
+ <canvas id="canvas" width="256" height="256"></canvas>
49
+ <div class="placeholder" id="placeholder">Load the model to start</div>
50
+ </div>
51
+ <div class="side">
52
+ <dl class="timings" id="timings"></dl>
53
+ <button id="save" class="ghost" disabled>Save PNG</button>
54
+ <p class="ai" id="ai-note" hidden>AI-generated image. It carries an invisible watermark and provenance fields in the saved PNG.</p>
55
+ </div>
56
+ </section>
57
+
58
+ <section class="about">
59
+ <h2>What this is</h2>
60
+ <p><a href="https://huggingface.co/ysharma/agate-002-4step" target="_blank" rel="noopener">ysharma/agate-002-4step</a> is an
61
+ <b>unofficial</b> distillation of <a href="https://huggingface.co/Logolabs/agate-preview-002" target="_blank" rel="noopener">Agate Preview 002</a> by LogoLabs.
62
+ Classifier-free guidance was baked into the weights, then the 50-step sampler was halved four times (32 → 16 → 8 → 4).
63
+ Each image is 4 network passes instead of Agate's 100. On the official GenEval scorer it gets 0.527, against 0.562 for the teacher at 50 steps and 0.466 for the teacher run at 4 steps.</p>
64
+ <p>256 × 256 only. No negative prompts (guidance is baked in). It inherits Agate's weaknesses: exact text, counts above three, negation.
65
+ There is no safety filter in this page, and the base model can draw people partly unclothed without being asked.</p>
66
+ <p class="fine">Runtime adapted from LogoLabs' <a href="https://huggingface.co/spaces/Logolabs/agate-webgpu" target="_blank" rel="noopener">agate-webgpu</a> (MIT):
67
+ tokenizer, text encoder, sampler and TAESD decoder run through onnxruntime-web. Seeds use a JavaScript PRNG and do not match PyTorch seeds.
68
+ This page is not affiliated with LogoLabs. Licence: MIT.</p>
69
+ </section>
70
+ </main>
71
+ <script type="module" src="js/app.js"></script>
72
+ </body>
73
  </html>
js/agate.js ADDED
@@ -0,0 +1,390 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ // Agate in the browser: tokenizer (tokenizers.js) + ONNX graphs run by onnxruntime-web (WebGPU, WASM
2
+ // fallback) + an Euler flow sampler with CFG. Mirrors agate/pipeline.py (AgatePipeline) of
3
+ // Logolabs/agate-preview-001 / -002 (fcdm_t2, 256 px) and Logolabs/agate-preview-003 (fcdm_t2mr: 512 or 256 px,
4
+ // SD3 timestep shift, prompt pipeline + count code, see prompt.js). Every image is marked (marking.js).
5
+
6
+ import * as ort from "https://cdn.jsdelivr.net/npm/onnxruntime-web@1.30.0/dist/ort.webgpu.bundle.min.mjs";
7
+ import { Tokenizer } from "https://cdn.jsdelivr.net/npm/@huggingface/tokenizers@0.2.0/dist/tokenizers.mjs";
8
+ import { prepare, countVector } from "./prompt.js";
9
+ import { embedWatermark, marks } from "./marking.js";
10
+ export const MODEL_REPO = "ysharma/agate-002-4step";
11
+
12
+ export { ort };
13
+
14
+ // The model this page runs: the unofficial 4-step distillation of Agate Preview 002
15
+ // (ysharma/agate-002-4step, MIT). Its webgpu/ folder has the same layout and graph I/O as 002's original webgpu/
16
+ // build (generator.onnx: z, t, ctx, mask -> v, plan; batch and text length dynamic).
17
+ const REPO = MODEL_REPO;
18
+ export const VERSIONS = {
19
+ "4step": { remote: `https://huggingface.co/${REPO}/resolve/main/webgpu/`, dir: "", cache: "agate-002-4step-web", res: [256], release: "002",
20
+ note: "Guidance- and progressively-distilled from Agate Preview 002: 4 Euler steps, no CFG. 256 px." },
21
+ };
22
+ const countUrl = () => `https://huggingface.co/${REPO}/resolve/main/config.json`;
23
+ const FILES_T2 = ["tokenizer.json", "tokenizer_config.json", "text_encoder.onnx", "generator.onnx", "taesd_decoder.onnx"];
24
+
25
+ export async function webgpuStatus() {
26
+ if (!("gpu" in navigator)) return { ok: false, why: "This browser does not expose WebGPU (navigator.gpu is missing)." };
27
+ try {
28
+ const adapter = await navigator.gpu.requestAdapter({ powerPreference: "high-performance" });
29
+ if (!adapter) return { ok: false, why: "WebGPU is present but no GPU adapter is available." };
30
+ let name = "";
31
+ try { const info = adapter.info || (await adapter.requestAdapterInfo?.()); name = [info?.vendor, info?.architecture, info?.description].filter(Boolean).join(" "); } catch { /* optional */ }
32
+ return { ok: true, adapter: name, maxBuffer: adapter.limits?.maxBufferSize };
33
+ } catch (e) {
34
+ return { ok: false, why: `WebGPU adapter request failed: ${e.message || e}` };
35
+ }
36
+ }
37
+
38
+ // ---- downloads with progress + Cache Storage ---------------------------------------------------
39
+ async function openCache(name) {
40
+ try { return await caches.open(name); } catch { return null; } // file://, private mode, ...
41
+ }
42
+
43
+ async function fetchCached(cacheName, url, key, expectBytes, onBytes) {
44
+ const cache = await openCache(cacheName);
45
+ if (cache) {
46
+ const hit = await cache.match(key);
47
+ if (hit) { const buf = new Uint8Array(await hit.arrayBuffer()); onBytes(buf.byteLength, true); return { buf, cached: true }; }
48
+ }
49
+ const res = await fetch(url);
50
+ if (!res.ok) throw new Error(`${url}: HTTP ${res.status}`);
51
+ // manifest size first: a compressed response would report the encoded length
52
+ const total = expectBytes || Number(res.headers.get("content-length")) || 0;
53
+ const reader = res.body.getReader();
54
+ let buf = new Uint8Array(total || 1 << 20), n = 0;
55
+ for (;;) {
56
+ const { done, value } = await reader.read();
57
+ if (done) break;
58
+ 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; }
59
+ buf.set(value, n); n += value.byteLength; onBytes(value.byteLength, false);
60
+ }
61
+ buf = n === buf.byteLength ? buf : buf.slice(0, n);
62
+ if (cache) {
63
+ try { await cache.put(key, new Response(buf, { headers: { "content-type": "application/octet-stream", "content-length": String(n) } })); }
64
+ catch (e) { console.warn("cache put failed (quota?)", e); }
65
+ }
66
+ return { buf, cached: false };
67
+ }
68
+
69
+ export async function hasCachedModel(version = "001") {
70
+ try { const c = await caches.open(VERSIONS[version].cache); return (await c.keys()).length >= 5; } catch { return false; }
71
+ }
72
+
73
+ export async function clearCache() {
74
+ let ok = true;
75
+ for (const v of Object.values(VERSIONS)) { try { ok = (await caches.delete(v.cache)) && ok; } catch { ok = false; } }
76
+ return ok;
77
+ }
78
+
79
+ // The Hub counts a model download only on a request for the model repo's root config.json; the ONNX files
80
+ // come from this Space, so every model load reads the loaded release's config.json once -- the same per-load
81
+ // count a Python from_pretrained() produces, cached weights or not.
82
+ function countLoad(version) { fetch(countUrl(version), { cache: "no-store" }).catch(() => {}); }
83
+
84
+ // ---- the model ---------------------------------------------------------------------------------
85
+ export class Agate {
86
+ // base: the page's models folder (001); override: an explicit models base for every version (local layout)
87
+ constructor({ base = "./models/", override = null, version = "003", ep = "webgpu", variant = null, optLevel = "all" } = {}) {
88
+ this.optLevel = optLevel;
89
+ const slash = (u) => (u.endsWith("/") ? u : u + "/");
90
+ this.version = version;
91
+ this.spec = VERSIONS[version];
92
+ this.base = override ? slash(override) + this.spec.dir : (this.spec.remote || slash(base) + this.spec.dir);
93
+ this.ep = ep;
94
+ this.variant = variant; // e.g. "full16": fp16-compute generator (experimental, 001 only)
95
+ this.marks = marks(this.spec.release, REPO);
96
+ }
97
+
98
+ get mr() { return this.manifest?.arch === "fcdm_t2mr"; }
99
+
100
+ // the fast build (2026-09-29): one Euler step per run, fp16 compute, everything stays on the GPU between steps
101
+ get fast() { return !!this.manifest?.step; }
102
+
103
+ fileKey(f) { return new Request(`${this.base}${f}?sha256=${this.manifest.files[f].sha256}`); }
104
+
105
+ async getFile(f, onBytes = () => {}) {
106
+ const meta = this.manifest.files[f];
107
+ return (await fetchCached(this.spec.cache, this.base + f, this.fileKey(f), meta.bytes, onBytes)).buf;
108
+ }
109
+
110
+ // onProgress({loaded, total, file, phase}); resolution: 003 only (512 default, or 256)
111
+ async load(onProgress = () => {}, resolution = null) {
112
+ const manifest = await (await fetch(this.base + "manifest.json", { cache: "no-cache" })).json();
113
+ this.manifest = manifest;
114
+ const swap = (this.variant && manifest.variants?.[this.variant]) || {};
115
+ const real = (f) => swap[f] || f;
116
+ this.resolution = this.mr ? Number(resolution || manifest.resolution) : 256;
117
+ // 003: download everything (the two graphs are ~1 MB each and share one weights file), so a resolution
118
+ // switch needs no network
119
+ const need = this.mr || manifest.step ? Object.keys(manifest.files) : FILES_T2;
120
+ const total = need.reduce((s, f) => s + manifest.files[real(f)].bytes, 0);
121
+ let loaded = 0, fromCache = 0;
122
+ const bufs = {};
123
+ const t0 = performance.now();
124
+ for (const f of need) {
125
+ const meta = manifest.files[real(f)];
126
+ const { buf } = await fetchCached(this.spec.cache, this.base + real(f), this.fileKey(real(f)), meta.bytes, (b, c) => {
127
+ loaded += b; if (c) fromCache += b; onProgress({ phase: "download", file: f, loaded, total });
128
+ });
129
+ bufs[f] = buf;
130
+ }
131
+ this.stats = { downloadMB: total / 1e6, fromCacheMB: fromCache / 1e6, downloadMs: performance.now() - t0 };
132
+ countLoad(this.version);
133
+ // drop cached files of older builds of this release (keys carry the sha256)
134
+ try {
135
+ const cache = await openCache(this.spec.cache), keep = new Set(need.map((f) => new URL(this.fileKey(real(f)).url, location.href).href));
136
+ if (cache) for (const req of await cache.keys()) if (!keep.has(req.url)) await cache.delete(req);
137
+ } catch (e) { console.warn("cache cleanup", e); }
138
+ const dec = new TextDecoder();
139
+ this.tok = new Tokenizer(JSON.parse(dec.decode(bufs["tokenizer.json"])), JSON.parse(dec.decode(bufs["tokenizer_config.json"])));
140
+
141
+ if (this.ep === "wasm") {
142
+ ort.env.wasm.numThreads = self.crossOriginIsolated ? Math.min(8, navigator.hardwareConcurrency || 4) : 1;
143
+ }
144
+ const t1 = performance.now();
145
+ this.sessions = {};
146
+ for (const [name, f] of [["text", "text_encoder.onnx"], ["vae", "taesd_decoder.onnx"]]) {
147
+ onProgress({ phase: "init", file: f, loaded: total, total });
148
+ const ts = performance.now();
149
+ this.sessions[name] = await ort.InferenceSession.create(bufs[f], this.sessionOptions());
150
+ this.stats[`${name}SessionMs`] = performance.now() - ts;
151
+ bufs[f] = null;
152
+ }
153
+ onProgress({ phase: "init", file: "generator", loaded: total, total });
154
+ const ts = performance.now();
155
+ if (this.fast) {
156
+ this.weights = bufs[manifest.step.external_data]; // kept for a resolution switch (no re-read)
157
+ await this.createGenerator(this.resolution, bufs);
158
+ } else if (this.mr) {
159
+ this.weights = bufs["generator.weights"]; // kept for a resolution switch (no re-read)
160
+ await this.createGenerator(this.resolution, bufs);
161
+ } else {
162
+ this.sessions.gen = await ort.InferenceSession.create(bufs["generator.onnx"], this.sessionOptions());
163
+ }
164
+ this.stats.genSessionMs = performance.now() - ts;
165
+ this.stats.sessionMs = performance.now() - t1;
166
+ return this.stats;
167
+ }
168
+
169
+ sessionOptions(extra = {}) { return { executionProviders: [this.ep], graphOptimizationLevel: this.optLevel, ...extra }; }
170
+
171
+ // 003: the generator graph for one resolution; the weights file is shared by both graphs.
172
+ async createGenerator(res, bufs = {}) {
173
+ let graphName, ext, extra = {};
174
+ if (this.fast) {
175
+ graphName = this.manifest.step.graphs[String(res)];
176
+ ext = this.manifest.step.external_data;
177
+ extra = { preferredOutputLocation: "gpu-buffer" }; // z_next / x1 / plan stay on the GPU
178
+ } else {
179
+ const r = this.manifest.resolutions[String(res)];
180
+ graphName = r?.graph; ext = r?.external_data;
181
+ }
182
+ if (!graphName) throw new Error(`resolution ${res} not available`);
183
+ const graph = bufs[graphName] || await this.getFile(graphName);
184
+ const weights = this.weights || await this.getFile(ext);
185
+ if (this.sessions.gen) { try { await this.sessions.gen.release(); } catch { /* */ } this.sessions.gen = null; }
186
+ this.sessions.gen = await ort.InferenceSession.create(graph, this.sessionOptions({ externalData: [{ path: ext, data: weights }], ...extra }));
187
+ this.resolution = Number(res);
188
+ }
189
+
190
+ async setResolution(res) {
191
+ if (!this.mr || Number(res) === this.resolution) return 0;
192
+ const t0 = performance.now();
193
+ await this.createGenerator(res);
194
+ return performance.now() - t0;
195
+ }
196
+
197
+ async release() {
198
+ for (const s of Object.values(this.sessions || {})) { try { await s?.release(); } catch { /* */ } }
199
+ this.sessions = null; this.weights = null;
200
+ }
201
+
202
+ tokenize(text) {
203
+ const max = this.manifest.text_max_len;
204
+ let ids = this.tok.encode(text).ids; // [CLS] ... [SEP], as the HF tokenizer
205
+ if (ids.length > max) ids = ids.slice(0, max - 1).concat(ids[ids.length - 1]); // truncation=True
206
+ return ids;
207
+ }
208
+
209
+ async encode(text) {
210
+ const ids = this.tokenize(text);
211
+ const n = ids.length;
212
+ const feeds = {
213
+ input_ids: new ort.Tensor("int64", BigInt64Array.from(ids, BigInt), [1, n]),
214
+ attention_mask: new ort.Tensor("int64", new BigInt64Array(n).fill(1n), [1, n]),
215
+ };
216
+ const out = await this.sessions.text.run(feeds);
217
+ const h = out.last_hidden_state;
218
+ const data = h.data.slice(); // (1, n, 512) fp32
219
+ h.dispose?.();
220
+ return { data, n, ids };
221
+ }
222
+
223
+ // The prompt as the model sees it: 003 runs its prompt pipeline (prompt.js), 001/002 use the prompt as typed.
224
+ prepare(prompt, negative = "") {
225
+ const pp = this.manifest.prompt_pipeline;
226
+ if (!pp) return { text: prompt, negative };
227
+ const [text, neg] = prepare(prompt, negative, pp.normalize, pp.spell);
228
+ return { text, negative: neg };
229
+ }
230
+
231
+ static gaussianNoise(seed, count) {
232
+ // mulberry32 + Box-Muller. Seeds do NOT reproduce PyTorch's torch.randn stream.
233
+ let a = seed >>> 0;
234
+ 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; };
235
+ const out = new Float32Array(count);
236
+ for (let i = 0; i < count; i += 2) {
237
+ const u1 = Math.max(rnd(), 1e-12), u2 = rnd();
238
+ const r = Math.sqrt(-2 * Math.log(u1));
239
+ out[i] = r * Math.cos(2 * Math.PI * u2);
240
+ if (i + 1 < count) out[i + 1] = r * Math.sin(2 * Math.PI * u2);
241
+ }
242
+ return out;
243
+ }
244
+
245
+ // SD3's resolution shift in Agate's convention (t = 0 noise), as agate/pipeline.py shift_t
246
+ static shiftT(t, shift) {
247
+ if (shift === 1) return t;
248
+ const s = 1 - t;
249
+ return 1 - shift * s / (1 + (shift - 1) * s);
250
+ }
251
+
252
+ // -> { rgba (watermarked unless watermark=false), width, height, latent, timings, prepared }
253
+ get hasPlan() { return !!this.sessions?.gen?.outputNames?.includes("plan"); }
254
+
255
+ // GPU-resident sampler for the step graphs: z_next of one run is the z input of the next, ctx / mask / counts
256
+ // are uploaded once, and nothing is read back until the end. The preview (plan + x1) is requested only when
257
+ // onPreview is set, previewReady() says the viewer is free and no earlier readback is still in flight; its
258
+ // download is started but never awaited by the loop (frames are dropped instead).
259
+ async sampleFast({ z, grid, cfg, hw, C, L, D, ctx, mask, cnt, onStep, onPreview, previewReady, shouldStop, stepMs }) {
260
+ const dev = ort.env.webgpu?.device;
261
+ const bufs = [];
262
+ const gpuT = (data, dims) => {
263
+ if (!dev) return new ort.Tensor("float32", data, dims);
264
+ const buf = dev.createBuffer({ size: Math.ceil(data.byteLength / 16) * 16, usage: GPUBufferUsage.STORAGE | GPUBufferUsage.COPY_SRC | GPUBufferUsage.COPY_DST });
265
+ dev.queue.writeBuffer(buf, 0, data);
266
+ bufs.push(buf);
267
+ return ort.Tensor.fromGpuBuffer(buf, { dataType: "float32", dims });
268
+ };
269
+ const fixed = { ctx: gpuT(ctx, [2, L, D]), mask: gpuT(mask, [2, L]) };
270
+ if (this.sessions.gen.inputNames.includes("counts")) fixed.counts = gpuT(cnt || new Float32Array(2 * L), [2, L]);
271
+ const steps = grid.length - 1;
272
+ let zT = new ort.Tensor("float32", z, [1, C, hw, hw]), inflight = false;
273
+ const one = (v) => new ort.Tensor("float32", new Float32Array([v]), [1]);
274
+ const cfgT = one(cfg);
275
+ try {
276
+ for (let i = 0; i < steps; i++) {
277
+ if (shouldStop()) throw new Error("stopped");
278
+ const ts = performance.now();
279
+ const want = !!onPreview && this.hasPlan && !inflight && previewReady();
280
+ const out = await this.sessions.gen.run({ z: zT, t: one(grid[i]), dt: one(grid[i + 1] - grid[i]), cfg: cfgT, ...fixed },
281
+ want ? ["z_next", "x1", "plan"] : ["z_next"]);
282
+ if (i > 0) zT.dispose?.();
283
+ zT = out.z_next;
284
+ if (want) {
285
+ inflight = true;
286
+ const step = i + 1;
287
+ Promise.all([out.plan.getData(true), out.x1.getData(true)])
288
+ .then(([plan, x1]) => onPreview({ step, steps, plan, x1, hw }))
289
+ .catch((e) => console.warn("preview", e))
290
+ .finally(() => { inflight = false; });
291
+ }
292
+ stepMs.push(performance.now() - ts);
293
+ onStep(i + 1, steps, null);
294
+ if ((i & 3) === 3) await new Promise((r) => setTimeout(r, 0)); // let the page paint now and then
295
+ }
296
+ const zf = await zT.getData(true); // the only synchronising readback
297
+ return Float32Array.from(zf);
298
+ } finally {
299
+ for (const b of bufs) { try { b.destroy(); } catch { /* */ } }
300
+ }
301
+ }
302
+
303
+ // onPreview({ step, steps, plan, x1, hw }) after every step, when given and the graph outputs the plan: plan =
304
+ // the thinker output for the conditional branch (Float32Array 640 x 16 x 16), x1 = z + (1 - t) v (guided).
305
+ async generate({ prompt, negative = "", seed = 0, steps = 50, cfg = 3.0, noise = null, watermark = true, onStep = () => {}, onPreview = null, previewReady = () => true, shouldStop = () => false }) {
306
+ const C = this.manifest.latent_ch, D = this.manifest.ctx_dim;
307
+ const rc = this.mr ? this.manifest.resolutions[String(this.resolution)] : { latent_hw: this.manifest.latent_hw, shift: 1 };
308
+ const hw = rc.latent_hw, shift = Number(rc.shift || 1);
309
+ const per = C * hw * hw;
310
+ const T = { };
311
+ let t0 = performance.now();
312
+ const prep = this.prepare(prompt, negative);
313
+ // cfg == 1 means v = v_cond: a batch of 1, no negative prompt (the distilled student's mode, as the patched
314
+ // agate/pipeline.py in the model repo does)
315
+ const single = cfg === 1 && !this.fast && !this.mr;
316
+ const nb = single ? 1 : 2;
317
+ const c = await this.encode(prep.text), u = single ? { data: new Float32Array(0), n: 0 } : await this.encode(prep.negative);
318
+ T.textMs = performance.now() - t0;
319
+ const BUCKETS = this.manifest.buckets;
320
+ const need = Math.max(c.n, u.n);
321
+ const L = BUCKETS.find((b) => b >= need) ?? BUCKETS[BUCKETS.length - 1];
322
+ const ctx = new Float32Array(nb * L * D), mask = new Float32Array(nb * L);
323
+ ctx.set(c.data, 0); mask.fill(1, 0, c.n); // zero padding, as F.pad in the pipeline
324
+ if (!single) { ctx.set(u.data, L * D); mask.fill(1, L, L + u.n); }
325
+ const feeds = { ctx: new ort.Tensor("float32", ctx, [nb, L, D]), mask: new ort.Tensor("float32", mask, [nb, L]) };
326
+ let countsNonzero = 0;
327
+ if (this.mr) { // count code: conditional half only
328
+ const cnt = new Float32Array(2 * L);
329
+ if (this.manifest.prompt_pipeline?.count_code) cnt.set(countVector(this.tok, prep.text, c.ids, Math.min(L, c.n)), 0);
330
+ countsNonzero = cnt.reduce((s, v) => s + (v > 0), 0);
331
+ feeds.counts = new ort.Tensor("float32", cnt, [2, L]);
332
+ }
333
+
334
+ let z = noise ? Float32Array.from(noise) : Agate.gaussianNoise(seed, per);
335
+ const grid = Array.from({ length: steps + 1 }, (_, i) => Agate.shiftT(i / steps, shift));
336
+ t0 = performance.now();
337
+ const stepMs = [];
338
+ if (this.fast) {
339
+ z = await this.sampleFast({ z, grid, cfg, hw, C, L, D, ctx, mask, cnt: feeds.counts?.data, onStep, onPreview, previewReady, shouldStop, stepMs });
340
+ } else {
341
+ const zz = new Float32Array(nb * per), tt = new Float32Array(nb);
342
+ for (let i = 0; i < steps; i++) {
343
+ if (shouldStop()) throw new Error("stopped");
344
+ const ts = performance.now();
345
+ zz.set(z, 0); if (!single) zz.set(z, per);
346
+ tt.fill(grid[i]);
347
+ const want = onPreview && this.hasPlan ? ["v", "plan"] : ["v"];
348
+ const out = await this.sessions.gen.run({ z: new ort.Tensor("float32", zz, [nb, C, hw, hw]), t: new ort.Tensor("float32", tt, [nb]), ...feeds }, want);
349
+ const v = out.v.data;
350
+ const dt = grid[i + 1] - grid[i];
351
+ const zn = new Float32Array(per);
352
+ let x1 = null;
353
+ if (out.plan) x1 = new Float32Array(per);
354
+ for (let k = 0; k < per; k++) {
355
+ const vc = v[k], vu = single ? vc : v[per + k], g = vu + cfg * (vc - vu);
356
+ zn[k] = z[k] + dt * g;
357
+ if (x1) x1[k] = z[k] + (1 - grid[i]) * g;
358
+ }
359
+ out.v.dispose?.();
360
+ if (out.plan) { const plan = out.plan.data.slice(); out.plan.dispose?.(); onPreview({ step: i + 1, steps, plan, x1, hw }); }
361
+ z = zn;
362
+ stepMs.push(performance.now() - ts);
363
+ onStep(i + 1, steps, z);
364
+ await new Promise((r) => setTimeout(r, 0)); // let the progress bar paint
365
+ }
366
+ }
367
+ T.samplerMs = performance.now() - t0;
368
+ T.firstStepMs = stepMs[0];
369
+ T.stepMs = stepMs.length > 1 ? (T.samplerMs - stepMs[0]) / (stepMs.length - 1) : stepMs[0];
370
+
371
+ t0 = performance.now();
372
+ const img = await this.sessions.vae.run({ latent: new ort.Tensor("float32", z, [1, C, hw, hw]) });
373
+ const x = img.image.data, H = img.image.dims[2], W = img.image.dims[3];
374
+ const rgba = new Uint8ClampedArray(H * W * 4);
375
+ for (let p = 0; p < H * W; p++) {
376
+ for (let ch = 0; ch < 3; ch++) {
377
+ const val = Math.min(1, Math.max(-1, x[ch * H * W + p]));
378
+ rgba[p * 4 + ch] = Math.round((val + 1) * 127.5);
379
+ }
380
+ rgba[p * 4 + 3] = 255;
381
+ }
382
+ T.decodeMs = performance.now() - t0;
383
+ t0 = performance.now();
384
+ if (watermark) embedWatermark(rgba, W, H, this.marks.payload);
385
+ T.markMs = performance.now() - t0;
386
+ T.totalMs = T.textMs + T.samplerMs + T.decodeMs + T.markMs;
387
+ T.bucket = L;
388
+ return { rgba, width: W, height: H, latent: z, timings: T, prepared: { ...prep, countsNonzero }, watermarked: watermark };
389
+ }
390
+ }
js/app.js ADDED
@@ -0,0 +1,169 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import { Agate, VERSIONS, webgpuStatus, hasCachedModel } from "./agate.js";
2
+ import { addPngText } from "./marking.js";
3
+
4
+ const VERSION = "4step";
5
+ const STEPS = 4, CFG = 1.0; // the student's sampler: 4 uniform Euler steps, guidance baked in
6
+ const params = new URLSearchParams(location.search);
7
+ let ep = params.get("ep") || "webgpu";
8
+
9
+ const EXAMPLES = [
10
+ "a red cube on top of a blue sphere",
11
+ "a green teapot and a red cup on a table",
12
+ "a minimalist logo of a fox head, orange, flat design, white background",
13
+ "a cabin in a snowy forest at night, warm light in the windows",
14
+ "an astronaut riding a horse on the moon",
15
+ "a watercolor painting of a harbor with red sailboats",
16
+ ];
17
+
18
+ const $ = (id) => document.getElementById(id);
19
+ const ui = {
20
+ load: $("load"), loadrow: $("loadrow"), fill: $("bar-fill"), status: $("status"), statusR: $("status-r"),
21
+ prompt: $("prompt"), live: $("live"), seed: $("seed"), dice: $("dice"), go: $("go"),
22
+ canvas: $("canvas"), frame: $("frame"), placeholder: $("placeholder"), timings: $("timings"),
23
+ save: $("save"), aiNote: $("ai-note"),
24
+ };
25
+ const store = {
26
+ get(k) { try { return localStorage.getItem(k); } catch { return null; } },
27
+ set(k, v) { try { localStorage.setItem(k, v); } catch { /* storage blocked */ } },
28
+ };
29
+
30
+ let agate = null, ready = false, running = false, pending = false, count = 0;
31
+ window.__agate = { state: "idle" }; // read by automated browser tests
32
+
33
+ function setProgress(frac, left, right = "") {
34
+ ui.fill.style.width = `${Math.round(Math.max(0, Math.min(1, frac)) * 100)}%`;
35
+ ui.status.textContent = left;
36
+ ui.statusR.textContent = right;
37
+ }
38
+
39
+ function renderExamples() {
40
+ for (const p of EXAMPLES) {
41
+ const b = document.createElement("button");
42
+ b.className = "chip"; b.type = "button"; b.textContent = p;
43
+ b.onclick = () => { ui.prompt.value = p; request(); };
44
+ $("examples").appendChild(b);
45
+ }
46
+ }
47
+
48
+ function showTimings(T) {
49
+ const rows = [
50
+ ["image", `${Math.round(T.totalMs)} ms`, true],
51
+ ["text encoder", `${T.textMs.toFixed(0)} ms`],
52
+ [`sampler (${STEPS} steps)`, `${T.samplerMs.toFixed(0)} ms`],
53
+ ["per step", `${T.stepMs.toFixed(0)} ms`],
54
+ ["decode (TAESD)", `${T.decodeMs.toFixed(0)} ms`],
55
+ ["watermark", `${T.markMs.toFixed(0)} ms`],
56
+ ["backend", ep === "wasm" ? "CPU (WASM)" : "WebGPU"],
57
+ ["images this session", String(count)],
58
+ ];
59
+ ui.timings.replaceChildren(...rows.flatMap(([k, v, big]) => {
60
+ const dt = document.createElement("dt"), dd = document.createElement("dd");
61
+ dt.textContent = k; dd.textContent = v; if (big) dd.className = "big";
62
+ return [dt, dd];
63
+ }));
64
+ }
65
+
66
+ async function load() {
67
+ ui.load.disabled = true;
68
+ window.__agate.state = "loading";
69
+ const t0 = performance.now();
70
+ try {
71
+ agate = new Agate({ version: VERSION, ep });
72
+ await agate.load(({ phase, file, loaded, total }) => {
73
+ if (phase === "download") setProgress(loaded / total, `Downloading ${file}`, `${(loaded / 1e6).toFixed(0)} / ${(total / 1e6).toFixed(0)} MB`);
74
+ else setProgress(1, `Starting ${file}`);
75
+ });
76
+ // the first WebGPU run compiles shaders: pay it here, not on the first keystroke
77
+ setProgress(1, "Warming up the GPU");
78
+ await agate.generate({ prompt: "a photo", seed: 0, steps: STEPS, cfg: CFG, watermark: false });
79
+ ready = true;
80
+ ui.loadrow.hidden = true;
81
+ ui.go.disabled = false;
82
+ ui.placeholder.hidden = true;
83
+ const s = ((performance.now() - t0) / 1000).toFixed(1);
84
+ setProgress(0, `Ready (${ep === "wasm" ? "CPU" : "WebGPU"})`, `loaded in ${s} s`);
85
+ window.__agate = { state: "ready", loadMs: performance.now() - t0, ep };
86
+ request();
87
+ } catch (e) {
88
+ console.error(e);
89
+ ui.load.disabled = false;
90
+ setProgress(0, "Could not load the model", String(e.message || e));
91
+ window.__agate = { state: "error", error: String(e.message || e) };
92
+ }
93
+ }
94
+
95
+ // Coalescing runner: while one image is being made, keystrokes only mark the prompt as dirty; the newest prompt is
96
+ // drawn as soon as the GPU is free, so the page never queues stale work.
97
+ function request() {
98
+ if (!ready) return;
99
+ pending = true;
100
+ if (!running) run();
101
+ }
102
+
103
+ async function run() {
104
+ running = true;
105
+ ui.frame.classList.add("busy");
106
+ while (pending) {
107
+ pending = false;
108
+ const prompt = ui.prompt.value.trim();
109
+ if (!prompt) continue;
110
+ const seed = Math.max(0, Math.min(4294967295, Number(ui.seed.value) || 0));
111
+ try {
112
+ window.__agate.state = "generating";
113
+ const r = await agate.generate({ prompt, seed, steps: STEPS, cfg: CFG, watermark: true });
114
+ const ctx = ui.canvas.getContext("2d");
115
+ ui.canvas.width = r.width; ui.canvas.height = r.height;
116
+ ctx.putImageData(new ImageData(r.rgba, r.width, r.height), 0, 0);
117
+ count += 1;
118
+ showTimings(r.timings);
119
+ ui.save.disabled = false;
120
+ ui.aiNote.hidden = false;
121
+ window.__agate = { state: "done", timings: r.timings, count, prompt };
122
+ } catch (e) {
123
+ console.error(e);
124
+ setProgress(0, "Error", String(e.message || e));
125
+ window.__agate = { state: "error", error: String(e.message || e) };
126
+ }
127
+ }
128
+ ui.frame.classList.remove("busy");
129
+ running = false;
130
+ }
131
+
132
+ async function save() {
133
+ const blob = await new Promise((r) => ui.canvas.toBlob(r, "image/png"));
134
+ const bytes = addPngText(new Uint8Array(await blob.arrayBuffer()), agate.marks.info);
135
+ const a = document.createElement("a");
136
+ a.href = URL.createObjectURL(new Blob([bytes], { type: "image/png" }));
137
+ a.download = `agate-4step-${ui.seed.value}.png`;
138
+ a.click();
139
+ setTimeout(() => URL.revokeObjectURL(a.href), 1000);
140
+ }
141
+
142
+ async function init() {
143
+ renderExamples();
144
+ ui.live.checked = store.get("agate4-live") !== "0";
145
+ ui.live.onchange = () => store.set("agate4-live", ui.live.checked ? "1" : "0");
146
+ ui.prompt.addEventListener("input", () => { if (ui.live.checked) request(); });
147
+ ui.prompt.addEventListener("keydown", (e) => { if (e.key === "Enter" && (e.ctrlKey || e.metaKey)) request(); });
148
+ ui.seed.addEventListener("input", () => { if (ui.live.checked) request(); });
149
+ ui.dice.onclick = () => { ui.seed.value = String(Math.floor(Math.random() * 1e9)); request(); };
150
+ ui.go.onclick = request;
151
+ ui.load.onclick = load;
152
+ ui.save.onclick = save;
153
+
154
+ if (ep !== "wasm") {
155
+ const s = await webgpuStatus();
156
+ if (!s.ok) {
157
+ $("nogpu").hidden = false;
158
+ $("nogpu-why").textContent = s.why;
159
+ ui.load.disabled = true;
160
+ $("use-wasm").onclick = () => { ep = "wasm"; $("nogpu").hidden = true; ui.load.disabled = false; ui.load.textContent = "Load model on the CPU (about 615 MB)"; };
161
+ window.__agate = { state: "no-webgpu", why: s.why };
162
+ return;
163
+ }
164
+ ui.statusR.textContent = s.adapter || "";
165
+ }
166
+ if (params.has("autoload") || (await hasCachedModel(VERSION))) load();
167
+ }
168
+
169
+ init();
js/marking.js ADDED
@@ -0,0 +1,188 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ // AI-generated content marking for the images this page makes (EU AI Act Art. 50(2)), the same marks as the
2
+ // Python packages (agate/marking.py of Logolabs/agate-preview-002 / -003):
3
+ // * an invisible watermark in the pixels: invisible-watermark's (MIT) 'dwtDctSvd' method, fixed 64-bit payload
4
+ // "AGATE" + release ("AGATE001" / "AGATE002" / "AGATE003"), ported from the library with its exact conventions
5
+ // (OpenCV 8-bit YUV, U channel only, Haar LL band, 4x4 DCT + SVD, s0 quantised to 36, the library's swapped
6
+ // detail order in the inverse DWT, numpy's truncating uint8 cast). agate.detect_watermark() in Python reads it.
7
+ // * provenance text chunks (tEXt) in the downloaded PNG: ai_generated, generator, model, watermark -- the keys
8
+ // and values the Python packages write. The prompt is NOT written.
9
+ // Neither mark is tamper-proof; see the model cards.
10
+
11
+ export const METHOD = "dwtDctSvd";
12
+ const SCALE = 36, BLOCK = 4;
13
+ const S2 = 0.7071067811865476; // PyWavelets' Haar tap, 1/sqrt(2)
14
+
15
+ // payload keeps the base release (AGATE002), as agate/marking.py in ysharma/agate-002-4step does; the provenance
16
+ // fields name the distilled model
17
+ export function marks(release, repo = `Logolabs/agate-preview-${release}`) {
18
+ const payload = "AGATE" + release; // 8 ASCII bytes = 64 bits
19
+ return {
20
+ payload,
21
+ info: {
22
+ ai_generated: "true",
23
+ generator: `Agate 002 4-step, unofficial distillation of Agate Preview ${release} (LogoLabs)`,
24
+ model: repo,
25
+ watermark: `invisible-watermark ${METHOD}, payload ${payload}`,
26
+ },
27
+ };
28
+ }
29
+
30
+ const desc = (x) => (x + 8192) >> 14; // OpenCV CV_DESCALE(x, 14)
31
+ const sat = (x) => (x < 0 ? 0 : x > 255 ? 255 : x);
32
+ const u8unsafe = (x) => { const t = Math.trunc(x) % 256; return t < 0 ? t + 256 : t; }; // numpy float -> uint8
33
+
34
+ // orthonormal 4-point DCT-II matrix (cv2.dct on a 4x4 block = C X C^T)
35
+ const C = [];
36
+ for (let k = 0; k < BLOCK; k++) {
37
+ C.push([]);
38
+ for (let n = 0; n < BLOCK; n++) C[k].push(Math.sqrt(k === 0 ? 1 / BLOCK : 2 / BLOCK) * Math.cos(Math.PI * (2 * n + 1) * k / (2 * BLOCK)));
39
+ }
40
+ function mul(A, B) { const R = [[0, 0, 0, 0], [0, 0, 0, 0], [0, 0, 0, 0], [0, 0, 0, 0]]; for (let i = 0; i < 4; i++) for (let j = 0; j < 4; j++) { let s = 0; for (let k = 0; k < 4; k++) s += A[i][k] * B[k][j]; R[i][j] = s; } return R; }
41
+ const T = (A) => A[0].map((_, j) => A.map((r) => r[j]));
42
+ const CT = T(C);
43
+
44
+ // largest singular value of a 4x4 matrix and its singular vectors (Jacobi on D^T D)
45
+ function topSVD(D) {
46
+ const M = mul(T(D), D);
47
+ const V = [[1, 0, 0, 0], [0, 1, 0, 0], [0, 0, 1, 0], [0, 0, 0, 1]];
48
+ for (let sweep = 0; sweep < 30; sweep++) {
49
+ let off = 0;
50
+ for (let p = 0; p < 4; p++) for (let q = p + 1; q < 4; q++) off += M[p][q] * M[p][q];
51
+ if (off < 1e-30) break;
52
+ for (let p = 0; p < 4; p++) for (let q = p + 1; q < 4; q++) {
53
+ if (Math.abs(M[p][q]) < 1e-300) continue;
54
+ const th = (M[q][q] - M[p][p]) / (2 * M[p][q]);
55
+ const t = Math.sign(th || 1) / (Math.abs(th) + Math.sqrt(th * th + 1));
56
+ const c = 1 / Math.sqrt(t * t + 1), s = t * c;
57
+ for (let k = 0; k < 4; k++) { const a = M[k][p], b = M[k][q]; M[k][p] = c * a - s * b; M[k][q] = s * a + c * b; }
58
+ for (let k = 0; k < 4; k++) { const a = M[p][k], b = M[q][k]; M[p][k] = c * a - s * b; M[q][k] = s * a + c * b; }
59
+ for (let k = 0; k < 4; k++) { const a = V[k][p], b = V[k][q]; V[k][p] = c * a - s * b; V[k][q] = s * a + c * b; }
60
+ }
61
+ }
62
+ let best = 0;
63
+ for (let i = 1; i < 4; i++) if (M[i][i] > M[best][best]) best = i;
64
+ const s0 = Math.sqrt(Math.max(0, M[best][best]));
65
+ let v = V.map((r) => r[best]), u;
66
+ if (s0 < 1e-12) { u = [1, 0, 0, 0]; v = [1, 0, 0, 0]; }
67
+ else u = D.map((r) => (r[0] * v[0] + r[1] * v[1] + r[2] * v[2] + r[3] * v[3]) / s0);
68
+ return { s0, u, v };
69
+ }
70
+
71
+ function payloadBits(payload) {
72
+ const bytes = new TextEncoder().encode(payload), bits = [];
73
+ for (const b of bytes) for (let k = 7; k >= 0; k--) bits.push((b >> k) & 1);
74
+ return bits;
75
+ }
76
+
77
+ // U channel (OpenCV BGR2YUV, 8 bit) of an RGBA buffer -> Int32Array(W*H), plus Y and V for the inverse
78
+ function toYUV(rgba, W, H) {
79
+ const Y = new Int32Array(W * H), U = new Int32Array(W * H), V = new Int32Array(W * H);
80
+ for (let p = 0; p < W * H; p++) {
81
+ const r = rgba[4 * p], g = rgba[4 * p + 1], b = rgba[4 * p + 2];
82
+ const y = desc(b * 1868 + g * 9617 + r * 4899);
83
+ Y[p] = y; U[p] = sat(desc((b - y) * 8061 + (128 << 14))); V[p] = sat(desc((r - y) * 14369 + (128 << 14)));
84
+ }
85
+ return { Y, U, V };
86
+ }
87
+
88
+ // LL band of the Haar DWT of U over the (H//4*4, W//4*4) region; -> {ca, ch, cv, cd} as Float64Array (h2 x w2)
89
+ function haar(U, W, R, Cc) {
90
+ const h2 = R / 2, w2 = Cc / 2, n = h2 * w2;
91
+ const ca = new Float64Array(n), ch = new Float64Array(n), cv = new Float64Array(n), cd = new Float64Array(n);
92
+ for (let i = 0; i < h2; i++) for (let j = 0; j < w2; j++) {
93
+ const a = U[(2 * i) * W + 2 * j], b = U[(2 * i) * W + 2 * j + 1], c = U[(2 * i + 1) * W + 2 * j], d = U[(2 * i + 1) * W + 2 * j + 1];
94
+ const k = i * w2 + j;
95
+ // PyWavelets' exact float order: pairs along rows first (a|c, b|d), then along columns
96
+ const l0 = a * S2 + c * S2, l1 = b * S2 + d * S2, h0 = a * S2 - c * S2, h1 = b * S2 - d * S2;
97
+ ca[k] = l0 * S2 + l1 * S2; cv[k] = l0 * S2 - l1 * S2; ch[k] = h0 * S2 + h1 * S2; cd[k] = h0 * S2 - h1 * S2;
98
+ }
99
+ return { ca, ch, cv, cd, h2, w2 };
100
+ }
101
+
102
+ function blockAt(ca, w2, bi, bj) { const B = []; for (let r = 0; r < 4; r++) { B.push([]); for (let c = 0; c < 4; c++) B[r].push(ca[(bi * 4 + r) * w2 + bj * 4 + c]); } return B; }
103
+
104
+ // Embed `payload` in place into rgba (Uint8ClampedArray / Uint8Array, W*H*4). Needs W*H >= 256*256.
105
+ export function embedWatermark(rgba, W, H, payload) {
106
+ if (W * H < 256 * 256) throw new Error("watermark: image too small (needs at least 256 x 256)");
107
+ const bits = payloadBits(payload);
108
+ const { Y, U, V } = toYUV(rgba, W, H);
109
+ const R = Math.floor(H / 4) * 4, Cc = Math.floor(W / 4) * 4;
110
+ const { ca, ch, cv, cd, h2, w2 } = haar(U, W, R, Cc);
111
+ const nbi = Math.floor(h2 / 4), nbj = Math.floor(w2 / 4);
112
+ let num = 0;
113
+ for (let bi = 0; bi < nbi; bi++) for (let bj = 0; bj < nbj; bj++, num++) {
114
+ const D = mul(mul(C, blockAt(ca, w2, bi, bj)), CT);
115
+ const { s0, u, v } = topSVD(D);
116
+ const s1 = (Math.floor(s0 / SCALE) + 0.25 + 0.5 * bits[num % bits.length]) * SCALE;
117
+ const Dn = D.map((row, i) => row.map((x, j) => x + (s1 - s0) * u[i] * v[j]));
118
+ const Bn = mul(mul(CT, Dn), C);
119
+ for (let r = 0; r < 4; r++) for (let c = 0; c < 4; c++) ca[(bi * 4 + r) * w2 + bj * 4 + c] = Bn[r][c];
120
+ }
121
+ // inverse Haar with the library's swapped details (cv as H, ch as V), truncating uint8 cast
122
+ for (let i = 0; i < h2; i++) for (let j = 0; j < w2; j++) {
123
+ // pywt.idwt2((ca, (cv, ch, cd))): the library's swapped details; PyWavelets' float order (columns, then rows)
124
+ const k = i * w2 + j, A = ca[k], hIn = cv[k], vIn = ch[k], d = cd[k];
125
+ const L0 = A * S2 + vIn * S2, L1 = A * S2 - vIn * S2, H0 = hIn * S2 + d * S2, H1 = hIn * S2 - d * S2;
126
+ U[(2 * i) * W + 2 * j] = u8unsafe(L0 * S2 + H0 * S2);
127
+ U[(2 * i) * W + 2 * j + 1] = u8unsafe(L1 * S2 + H1 * S2);
128
+ U[(2 * i + 1) * W + 2 * j] = u8unsafe(L0 * S2 - H0 * S2);
129
+ U[(2 * i + 1) * W + 2 * j + 1] = u8unsafe(L1 * S2 - H1 * S2);
130
+ }
131
+ for (let p = 0; p < W * H; p++) { // OpenCV YUV2BGR, 8 bit
132
+ const y = Y[p], uu = U[p] - 128, vq = V[p] - 128;
133
+ rgba[4 * p + 2] = sat(y + desc(uu * 33292));
134
+ rgba[4 * p + 1] = sat(y + desc(uu * -6472 + vq * -9519));
135
+ rgba[4 * p] = sat(y + desc(vq * 18678));
136
+ }
137
+ return rgba;
138
+ }
139
+
140
+ // -> {bits, text, bitAccuracy} against `payload` (a self-check; the reference detector is the Python one)
141
+ export function readWatermark(rgba, W, H, payload) {
142
+ const want = payloadBits(payload), n = want.length;
143
+ const { U } = toYUV(rgba, W, H);
144
+ const R = Math.floor(H / 4) * 4, Cc = Math.floor(W / 4) * 4;
145
+ const { ca, h2, w2 } = haar(U, W, R, Cc);
146
+ const nbi = Math.floor(h2 / 4), nbj = Math.floor(w2 / 4);
147
+ const sum = new Float64Array(n), cnt = new Float64Array(n);
148
+ let num = 0;
149
+ for (let bi = 0; bi < nbi; bi++) for (let bj = 0; bj < nbj; bj++, num++) {
150
+ const { s0 } = topSVD(mul(mul(C, blockAt(ca, w2, bi, bj)), CT));
151
+ sum[num % n] += (s0 % SCALE) > SCALE * 0.5 ? 1 : 0; cnt[num % n] += 1;
152
+ }
153
+ const bits = Array.from(sum, (s, k) => (s / cnt[k]) * 255 > 127 ? 1 : 0);
154
+ let same = 0; bits.forEach((b, k) => { if (b === want[k]) same++; });
155
+ const bytes = []; for (let k = 0; k < n; k += 8) { let b = 0; for (let q = 0; q < 8; q++) b = (b << 1) | bits[k + q]; bytes.push(b); }
156
+ return { bits, text: String.fromCharCode(...bytes), bitAccuracy: same / n };
157
+ }
158
+
159
+ // ---- PNG tEXt chunks ---------------------------------------------------------------------------
160
+ let CRC_TABLE = null;
161
+ function crc32(bytes) {
162
+ if (!CRC_TABLE) { CRC_TABLE = new Uint32Array(256); for (let n = 0; n < 256; n++) { let c = n; for (let k = 0; k < 8; k++) c = c & 1 ? 0xedb88320 ^ (c >>> 1) : c >>> 1; CRC_TABLE[n] = c >>> 0; } }
163
+ let c = 0xffffffff;
164
+ for (const b of bytes) c = CRC_TABLE[(c ^ b) & 0xff] ^ (c >>> 8);
165
+ return (c ^ 0xffffffff) >>> 0;
166
+ }
167
+
168
+ // PNG bytes -> PNG bytes with one tEXt chunk per entry of `info` (Latin-1 keys/values), inserted before IEND.
169
+ export function addPngText(png, info) {
170
+ const latin1 = (s) => Uint8Array.from(s, (ch) => { const c = ch.charCodeAt(0); return c < 256 ? c : 63; });
171
+ const chunks = [];
172
+ for (const [k, v] of Object.entries(info)) {
173
+ const data = new Uint8Array([...latin1(k), 0, ...latin1(String(v))]);
174
+ const typeData = new Uint8Array([116, 69, 88, 116, ...data]); // "tEXt"
175
+ const out = new Uint8Array(12 + data.length), dv = new DataView(out.buffer);
176
+ dv.setUint32(0, data.length); out.set(typeData, 4); dv.setUint32(8 + data.length, crc32(typeData));
177
+ chunks.push(out);
178
+ }
179
+ // find IEND (the last chunk): 12 bytes from the end in every well-formed PNG
180
+ const iend = png.length - 12;
181
+ const extra = chunks.reduce((s, c) => s + c.length, 0);
182
+ const res = new Uint8Array(png.length + extra);
183
+ res.set(png.subarray(0, iend), 0);
184
+ let o = iend;
185
+ for (const c of chunks) { res.set(c, o); o += c.length; }
186
+ res.set(png.subarray(iend), o);
187
+ return res;
188
+ }
js/prompt.js ADDED
@@ -0,0 +1,249 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ // Agate Preview 003 prompt pipeline, a line-by-line port of agate/prompt_norm.py (+ AgatePipeline.prepare
2
+ // and prompt_norm.count_tensor) of Logolabs/agate-preview-003. Checked against the Python package by
3
+ // test/prompt_golden.json (tools: export/prompt_golden.py writes it, test/prompt_golden.mjs or ?golden=1 checks it).
4
+ //
5
+ // Python-compatibility notes (the reason for the helpers below):
6
+ // * Python's \s, \d, \b and str.isalnum/isspace are Unicode-aware; JS's are ASCII. PY_WS, \p{Nd} and the
7
+ // Unicode word-boundary WB mirror Python's definitions (\w = letters, numbers, underscore).
8
+ // * Python counts code points; JS strings count UTF-16 units. Lengths that gate behaviour use cpLen().
9
+ // * Token offsets come from the HF fast tokenizer in Python; here they are rebuilt from the byte-level BPE
10
+ // tokens (tokenOffsets), in UTF-16 units, which gives the same overlaps as Python's code-point offsets.
11
+
12
+ const PY_WS = "\\t\\n\\x0b\\x0c\\r\\x1c-\\x20\\x85\\xa0\\u1680\\u2000-\\u200a\\u2028\\u2029\\u202f\\u205f\\u3000";
13
+ const S = `[${PY_WS}]`; // Python's \s
14
+ const NS = `[^${PY_WS}]`; // Python's \S
15
+ const W = "[\\p{L}\\p{N}_]"; // Python's \w (str patterns)
16
+ const WB = `(?:(?<=${W})(?!${W})|(?<!${W})(?=${W}))`; // Python's \b
17
+ const D = "\\p{Nd}"; // Python's \d
18
+ const re = (src, flags = "") => new RegExp(src, flags.includes("u") ? flags : flags + "u");
19
+
20
+ const NUM_LIST = "zero one two three four five six seven eight nine ten eleven twelve thirteen fourteen fifteen sixteen seventeen eighteen nineteen twenty".split(" ");
21
+ const NUM_WORDS = new Map(NUM_LIST.map((w, i) => [w, i]));
22
+ NUM_WORDS.set("dozen", 12);
23
+ const WORD_OF = new Map(NUM_LIST.map((w, i) => [i, w]));
24
+ const NOT_COUNT_NEXT = new Set(["pm", "am", "o'clock", "percent", "%", "years", "year", "hours", "hour", "minutes", "minute",
25
+ "seconds", "second", "times", "px", "cm", "mm", "m", "km", "kg", "g", "inch", "inches", "feet",
26
+ "degrees", "x", "d", "k", "th", "st", "nd", "rd", "of", "o"]);
27
+ const ONE_PRONOUN_PREV = new Set(["the", "this", "that", "no", "each", "every", "any", "some", "which", "another", "someone",
28
+ "everyone", "anyone", "only"]);
29
+ export const SPELL_SEP = " || spell: ";
30
+ const SPELL_MAX_FRAC = 0.6, SPELL_MAX_CHARS = 40, SPELL_MAX_WORDS = 6, SPELL_MAX_SPANS = 3;
31
+
32
+ const QUOTE_SRC = '"[^"]*"|“[^”]*”';
33
+ const cpLen = (s) => { let n = 0; for (const _ of s) n++; return n; }; // len() in Python
34
+ const isSpace = (ch) => re(`^${S}$`).test(ch);
35
+ const pyStrip = (s, chars = null) => {
36
+ const f = chars === null ? isSpace : (c) => chars.includes(c);
37
+ let a = 0, b = s.length;
38
+ while (a < b && f(s[a])) a++;
39
+ while (b > a && f(s[b - 1])) b--;
40
+ return s.slice(a, b);
41
+ };
42
+ const pySplit = (s) => s.split(re(`${S}+`)).filter((x) => x.length > 0); // str.split()
43
+ const isAlnum = (ch) => /^[\p{L}\p{N}]$/u.test(ch);
44
+ const isAlpha = (w) => w.length > 0 && /^\p{L}+$/u.test(w);
45
+ const isUpper = (w) => w !== w.toLowerCase() && w === w.toUpperCase(); // ASCII words only reach these
46
+ const isLower = (w) => w !== w.toUpperCase() && w === w.toLowerCase();
47
+ const isDigits = (w) => /^\p{Nd}+$/u.test(w);
48
+
49
+ function digitValue(cp) { // value of one \p{Nd} code point: Nd code points come in runs of ten
50
+ let start = cp;
51
+ while (/\p{Nd}/u.test(String.fromCodePoint(start - 1))) start--;
52
+ return (cp - start) % 10;
53
+ }
54
+ function pyInt(w) { // int() of a \p{Nd}+ string
55
+ let n = 0;
56
+ for (const ch of w) n = n * 10 + digitValue(ch.codePointAt(0));
57
+ return n;
58
+ }
59
+
60
+ function outsideQuotes(text) {
61
+ const spans = [];
62
+ let last = 0;
63
+ for (const m of text.matchAll(re(QUOTE_SRC, "g"))) { spans.push([last, m.index]); last = m.index + m[0].length; }
64
+ spans.push([last, text.length]);
65
+ return spans.filter(([a, b]) => b > a);
66
+ }
67
+
68
+ function mapOutside(text, fn) {
69
+ let out = "", last = 0;
70
+ for (const [a, b] of outsideQuotes(text)) { out += text.slice(last, a) + fn(text.slice(a, b)); last = b; }
71
+ return out + text.slice(last);
72
+ }
73
+
74
+ function fixCase(words) {
75
+ const long2 = words.filter((w) => cpLen(w) >= 2 && isAlpha(w));
76
+ if (long2.length && long2.filter(isUpper).length / long2.length >= 0.6) return "shout";
77
+ const long3 = words.filter((w) => cpLen(w) >= 3 && isAlpha(w));
78
+ if (long3.length >= 3) {
79
+ const t = long3.filter((w) => { const c = [...w]; return isUpper(c[0]) && isLower(c.slice(1).join("")); }).length;
80
+ if (t / long3.length >= 0.7) return "title";
81
+ }
82
+ return null;
83
+ }
84
+
85
+ export function isJsonCaption(text) {
86
+ const s = pyStrip(text.split(SPELL_SEP)[0]);
87
+ if (!(s.startsWith("{") && s.endsWith("}"))) return false;
88
+ try { const v = JSON.parse(s); return v !== null && typeof v === "object" && !Array.isArray(v); } catch { return false; }
89
+ }
90
+
91
+ function isCountContext(seg, end) {
92
+ const nxt = re(`^${S}*([A-Za-z%']+)`).exec(seg.slice(end));
93
+ return !!nxt && !NOT_COUNT_NEXT.has(nxt[1].toLowerCase());
94
+ }
95
+
96
+ function canonNumbers(seg) {
97
+ return seg.replace(re(`${WB}[A-Za-z]+${WB}|${WB}${D}+${WB}`, "g"), (w, off) => {
98
+ const lw = w.toLowerCase();
99
+ if (NUM_WORDS.has(lw)) return lw;
100
+ if (isDigits(w)) { const n = pyInt(w); if (n >= 1 && n <= 20 && isCountContext(seg, off + w.length)) return WORD_OF.get(n); }
101
+ return w;
102
+ });
103
+ }
104
+
105
+ export function normalize(text) {
106
+ if (isJsonCaption(text)) return text;
107
+ text = pyStrip(text.replace(re(`${S}+`, "g"), " "));
108
+ const words = [];
109
+ for (const [a, b] of outsideQuotes(text)) for (const m of text.slice(a, b).matchAll(re(`[A-Za-z][A-Za-z'\\-]*|${D}+`, "g"))) words.push(m[0]);
110
+ const mode = fixCase(words);
111
+ if (mode === "shout") text = mapOutside(text, (s) => s.toLowerCase());
112
+ else if (mode === "title") {
113
+ const first = re(`^${S}*${NS}+`).exec(text);
114
+ const head = first ? first[0].length : 0;
115
+ text = text.slice(0, head) + mapOutside(text.slice(head), (s) => s.replace(re(`${WB}[A-Z][a-z]*${WB}`, "g"), (m) => m.toLowerCase()));
116
+ }
117
+ return mapOutside(text, canonNumbers);
118
+ }
119
+
120
+ export function findCounts(text) {
121
+ const out = [];
122
+ if (isJsonCaption(text)) return out;
123
+ text = text.split(SPELL_SEP)[0];
124
+ for (const [a, b] of outsideQuotes(text)) {
125
+ const seg = text.slice(a, b);
126
+ const pat = re(`${WB}(?:a${S}+)?(dozen|pair(?=${S}+of${WB}))${WB}|${WB}([A-Za-z]+|${D}+)${WB}`, "gd");
127
+ for (const m of seg.matchAll(pat)) {
128
+ if (m[1] !== undefined) {
129
+ const [s1, e1] = m.indices[1];
130
+ out.push([a + s1, a + e1, m[1].toLowerCase() === "dozen" ? 12 : 2]);
131
+ continue;
132
+ }
133
+ const w = m[2], lw = w.toLowerCase();
134
+ const [s2, e2] = m.indices[2];
135
+ let n;
136
+ if (isDigits(lw)) { n = pyInt(lw); if (!(n >= 1 && n <= 99)) continue; }
137
+ else if (NUM_WORDS.has(lw) && lw !== "zero" && lw !== "dozen") n = NUM_WORDS.get(lw);
138
+ else continue;
139
+ if (!isCountContext(seg, m.index + m[0].length)) continue;
140
+ const prev = seg.slice(0, m.index).match(/[A-Za-z]+/g) || [];
141
+ if (lw === "one" && prev.length && ONE_PRONOUN_PREV.has(prev[prev.length - 1].toLowerCase())) continue;
142
+ out.push([a + s2, a + e2, n]);
143
+ }
144
+ }
145
+ return out;
146
+ }
147
+
148
+ export function textSpans(text) {
149
+ const body = text.split(SPELL_SEP)[0];
150
+ const out = [];
151
+ for (const m of body.matchAll(re(QUOTE_SRC, "g"))) {
152
+ const c = pyStrip(m[0].slice(1, -1));
153
+ if (c && [...c].some(isAlnum) && cpLen(c) <= SPELL_MAX_FRAC * cpLen(body) && cpLen(c) <= SPELL_MAX_CHARS && pySplit(c).length <= SPELL_MAX_WORDS) out.push(c);
154
+ }
155
+ return out.slice(0, SPELL_MAX_SPANS);
156
+ }
157
+
158
+ export function spell(content) {
159
+ const words = pySplit(content).map((w) => [...w].filter((ch) => isAlnum(ch) || "&!?'-".includes(ch)));
160
+ return words.filter((w) => w.length).map((w) => w.join(" ")).join(" / ");
161
+ }
162
+
163
+ export function addSpelling(text) {
164
+ if (text.includes(SPELL_SEP) || isJsonCaption(text)) return text;
165
+ const spans = textSpans(text);
166
+ if (!spans.length) return text;
167
+ return text + SPELL_SEP + spans.map(spell).join(" ; ");
168
+ }
169
+
170
+ export function splitNegatives(text) {
171
+ const negs = [];
172
+ if (isJsonCaption(text)) return [text, negs];
173
+ const pat = re(`${S}*,?${S}*${WB}(?:with${S}+)?(?:without|no)${S}+(?:any${S}+)?(?:a${S}+|an${S}+|the${S}+)?([^,.;"“]+?)(?=${S}+and${S}+|[,.;]|$)`, "gi");
174
+ const cut = (seg) => seg.replace(pat, (...g) => { negs.push(pyStrip(g[1])); return ""; });
175
+ let pos = mapOutside(text, cut);
176
+ pos = pyStrip(pos.replace(re(`${S}+([,.;])`, "g"), "$1").replace(re(`${S}+`, "g"), " "), " ,;");
177
+ return [pos, negs];
178
+ }
179
+
180
+ // AgatePipeline.prepare: -> [prompt as the text encoder sees it, negative prompt]
181
+ export function prepare(prompt, negativePrompt = "", normalizeOn = true, spellOn = true) {
182
+ let negs = [];
183
+ if (normalizeOn) [prompt, negs] = splitNegatives(normalize(prompt));
184
+ if (spellOn) prompt = addSpelling(prompt);
185
+ const negative = [pyStrip(negativePrompt), ...negs].filter((n) => n).join(", ");
186
+ return [prompt, negative];
187
+ }
188
+
189
+ // ---- token offsets for a byte-level BPE tokenizer (tokenizers.js has no offset mapping) -----------
190
+ let BYTE_OF = null; // GPT-2 bytes_to_unicode, inverted: char -> byte
191
+ function byteDecoder() {
192
+ if (BYTE_OF) return BYTE_OF;
193
+ const bs = [];
194
+ for (let b = 33; b <= 126; b++) bs.push(b);
195
+ for (let b = 161; b <= 172; b++) bs.push(b);
196
+ for (let b = 174; b <= 255; b++) bs.push(b);
197
+ const cs = bs.slice();
198
+ let n = 0;
199
+ for (let b = 0; b < 256; b++) if (!bs.includes(b)) { bs.push(b); cs.push(256 + n); n++; }
200
+ BYTE_OF = new Map(bs.map((b, i) => [String.fromCodePoint(cs[i]), b]));
201
+ return BYTE_OF;
202
+ }
203
+
204
+ // ids: the token ids of `text` ([CLS] ... [SEP]). -> [[start, end], ...] in UTF-16 units of text, [0, 0] for
205
+ // special tokens: Python's offset_mapping (code points) mapped to JS string indices.
206
+ export function tokenOffsets(tok, text, ids) {
207
+ const dec = byteDecoder();
208
+ const nfc = text.normalize("NFC");
209
+ const enc = new TextEncoder();
210
+ // byte index -> [UTF-16 start, UTF-16 end] of the character that byte belongs to
211
+ const charOfByte = [];
212
+ let u = 0;
213
+ for (const ch of nfc) { const nb = enc.encode(ch).length; for (let k = 0; k < nb; k++) charOfByte.push([u, u + ch.length]); u += ch.length; }
214
+ const added = tok.get_added_tokens_decoder?.() ?? new Map();
215
+ const out = [];
216
+ let pos = 0;
217
+ for (const id of ids) {
218
+ const at = added.get(Number(id));
219
+ let nbytes;
220
+ if (at && at.special) { out.push([0, 0]); continue; }
221
+ if (at) nbytes = enc.encode(at.content).length;
222
+ else {
223
+ const s = tok.id_to_token(Number(id));
224
+ nbytes = 0;
225
+ for (const ch of s) nbytes += dec.has(ch) ? 1 : enc.encode(ch).length;
226
+ }
227
+ if (nbytes === 0 || pos >= charOfByte.length) { out.push([0, 0]); continue; }
228
+ const a = charOfByte[pos][0], b = charOfByte[Math.min(pos + nbytes, charOfByte.length) - 1][1];
229
+ out.push([a, b]);
230
+ pos += nbytes;
231
+ }
232
+ // text that NFC changed: offsets refer to the NFC string; the count spans are found in the NFC string too
233
+ return out;
234
+ }
235
+
236
+ // prompt_norm.count_tensor for one prompt: Float32Array(length), the count value at the tokens that spell one.
237
+ export function countVector(tok, text, ids, length) {
238
+ const out = new Float32Array(length);
239
+ const nfc = text.normalize("NFC");
240
+ const spans = findCounts(nfc);
241
+ if (!spans.length) return out;
242
+ const offs = tokenOffsets(tok, nfc, ids);
243
+ for (let j = 0; j < Math.min(length, offs.length); j++) {
244
+ const [a, b] = offs[j];
245
+ if (b <= a) continue;
246
+ for (const [s, e, n] of spans) if (a < e && b > s) out[j] = n;
247
+ }
248
+ return out;
249
+ }