Agate 4-step WebGPU: in-browser runtime with live mode
Browse files- README.md +35 -5
- css/style.css +79 -0
- index.html +71 -17
- js/agate.js +390 -0
- js/app.js +169 -0
- js/marking.js +188 -0
- js/prompt.js +249 -0
README.md
CHANGED
|
@@ -1,10 +1,40 @@
|
|
| 1 |
---
|
| 2 |
-
title: Agate
|
| 3 |
-
emoji:
|
| 4 |
-
colorFrom:
|
| 5 |
-
colorTo:
|
| 6 |
sdk: static
|
|
|
|
| 7 |
pinned: false
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 8 |
---
|
| 9 |
|
| 10 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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 |
-
|
| 4 |
-
|
| 5 |
-
|
| 6 |
-
|
| 7 |
-
|
| 8 |
-
|
| 9 |
-
|
| 10 |
-
|
| 11 |
-
|
| 12 |
-
|
| 13 |
-
|
| 14 |
-
|
| 15 |
-
|
| 16 |
-
|
| 17 |
-
|
| 18 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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 |
+
}
|