File size: 11,099 Bytes
72ecd4e
 
 
2a40dc6
aedfd27
 
 
2a40dc6
 
 
aedfd27
72ecd4e
 
eb14405
2a40dc6
 
1ae0fdf
72ecd4e
2a40dc6
 
 
aedfd27
 
 
2a40dc6
 
aedfd27
 
 
 
 
72ecd4e
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
8a4f1e5
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
72ecd4e
 
 
 
 
 
 
 
 
 
 
 
 
8a4f1e5
 
 
72ecd4e
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
8095b6a
 
 
 
 
 
 
 
 
 
 
 
bc7639e
2a40dc6
 
1ae0fdf
bcb786b
1ae0fdf
 
72ecd4e
1ae0fdf
bc7639e
553984b
bc7639e
 
8095b6a
 
 
7aa9ecf
bcb786b
7aa9ecf
bcb786b
bc7639e
 
 
 
 
 
0cf86d0
553984b
bc7639e
 
 
 
553984b
 
 
 
 
 
 
 
72ecd4e
 
61a47e2
 
 
 
 
 
 
 
 
 
 
852cf61
7aa9ecf
72ecd4e
 
 
852cf61
72ecd4e
852cf61
72ecd4e
 
 
9d3fb9f
 
 
 
 
7aa9ecf
72ecd4e
 
 
 
 
 
 
 
 
0cf86d0
72ecd4e
 
7aa9ecf
852cf61
7aa9ecf
 
 
0cf86d0
 
 
 
 
 
72ecd4e
 
 
 
61a47e2
72ecd4e
 
 
 
 
 
 
 
 
 
 
 
 
 
61a47e2
852cf61
72ecd4e
852cf61
72ecd4e
 
 
61a47e2
 
72ecd4e
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1ae0fdf
72ecd4e
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
// Minimal browser inference for Bartholomheow/Supra2-IMG-ONNX — no build step.
// Serve this folder over http(s) and open index.html (modules + WebGPU need secure context).
//
// Needs the onnxruntime-web UMD global (plain <script> tag, pinned version):
//
// <script src="https://cdn.jsdelivr.net/npm/onnxruntime-web@1.30.0/dist/ort.all.min.js"></script>
//
// Transformers.js ships as ESM only (no UMD global exists), so it is imported
// below from the CDN by full URL. In a bundler (vite/webpack) both libraries
// resolve through ESM imports instead:
//   npm i onnxruntime-web @huggingface/transformers

const REPO = 'Bartholomheow/Supra2-IMG-ONNX';
export const EXAMPLE_BUILD = '2026-09-25e-fp32only';
const TRANSFORMERS_CDN_URL =
  'https://cdn.jsdelivr.net/npm/@huggingface/transformers@3.8.1/dist/transformers.min.js';
const fileUrl = (repo, path) => `https://huggingface.co/${repo}/resolve/main/${path}`;

// ESM import when bundled; UMD global (ort) or full-URL ESM import
// (transformers.js, which has no UMD build) when loaded via <script> tags.
async function loadLib(specifier, { globalName, url } = {}) {
  try {
    return await import(/* @vite-ignore */ specifier);
  } catch {
    if (url) return await import(/* @vite-ignore */ url);
    const g = globalName ? globalThis[globalName] : undefined;
    if (!g) throw new Error(`${specifier} failed to load (no bundle import and no window.${globalName})`);
    return g;
  }
}

function mulberry32(seed) {
  let a = seed >>> 0;
  return () => {
    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;
  };
}

function gaussian(seed, n) {
  const rand = mulberry32(seed);
  const out = new Float32Array(n);
  for (let i = 0; i < n; i += 2) {
    const r = Math.sqrt(-2 * Math.log(Math.max(rand(), 1e-12)));
    const a = 2 * Math.PI * rand();
    out[i] = r * Math.cos(a);
    if (i + 1 < n) out[i + 1] = r * Math.sin(a);
  }
  return out;
}

async function download(url, file, onProgress) {
  const cache = await caches.open('supra2-img-v1');
  const hit = await cache.match(url);
  if (hit) {
    const buf = await hit.arrayBuffer();
    // Self-healing: a poisoned (truncated) entry disagrees with the server
    // length — drop it and fall through to a fresh download.
    try {
      const head = await fetch(url, { method: 'HEAD' });
      const total = Number(head.headers.get('content-length')) || null;
      if (total == null || buf.byteLength === total) {
        onProgress?.({ file, loaded: buf.byteLength, total: total ?? buf.byteLength });
        return buf;
      }
      await cache.delete(url);
    } catch {
      onProgress?.({ file, loaded: buf.byteLength, total: buf.byteLength });
      return buf; // offline: serve what we have
    }
  }
  const res = await fetch(url);
  if (!res.ok) throw new Error(`download failed (${res.status}): ${file}`);
  const total = Number(res.headers.get('content-length')) || null;
  const reader = res.body.getReader();
  const chunks = [];
  let loaded = 0;
  for (;;) {
    const { done, value } = await reader.read();
    if (done) break;
    chunks.push(value);
    loaded += value.byteLength;
    onProgress?.({ file, loaded, total });
  }
  if (total != null && loaded !== total) {
    throw new Error(`truncated download: ${file} got ${loaded}/${total} bytes — retry`);
  }
  const buf = new Uint8Array(loaded);
  let off = 0;
  for (const c of chunks) { buf.set(c, off); off += c.byteLength; }
  await cache.put(url, new Response(buf.slice(0)));
  return buf.buffer;
}

function halfToFloat(h) {
  const s = (h & 0x8000) << 16;
  const e = (h >> 10) & 0x1f;
  const m = h & 0x3ff;
  const f = new Float32Array(1);
  const u = new Uint32Array(f.buffer);
  if (e === 0) {
    if (m === 0) { u[0] = s; return f[0]; }
    let mm = m, ee = -14;
    while ((mm & 0x400) === 0) { mm <<= 1; ee -= 1; }
    u[0] = s | ((ee + 127) << 23) | ((mm & 0x3ff) << 13);
    return f[0];
  }
  if (e === 31) { u[0] = s | 0x7f800000 | (m << 13); return f[0]; }
  u[0] = s | ((e + 112) << 23) | (m << 13);
  return f[0];
}

// onnxruntime silently falls back to WASM when WebGPU init fails, so a
// requested 'webgpu' backend proves nothing — check the adapter ourselves.
async function webgpuUsable() {
  try {
    const gpu = globalThis.navigator?.gpu;
    if (!gpu) return false;
    return (await gpu.requestAdapter()) !== null;
  } catch {
    return false;
  }
}

export async function loadSupra(repo = REPO, { onProgress, backends = ['webgpu'] } = {}) {
  const ort = await loadLib('onnxruntime-web', { globalName: 'ort' });
  const { AutoTokenizer } = await loadLib('@huggingface/transformers', { url: TRANSFORMERS_CDN_URL });
  const cfg = await (await fetch(fileUrl(repo, 'pipeline_config.json'))).json();
  const [ditBuf, vaeBuf] = await Promise.all([
    download(fileUrl(repo, cfg.dit), cfg.dit, onProgress),
    download(fileUrl(repo, cfg.vae_decoder), cfg.vae_decoder, onProgress),
  ]);
  const tokenize = await AutoTokenizer.from_pretrained(repo);
  const errors = [];
  let selected = null;
  for (const backend of backends) {
    try {
      if (backend === 'webgpu' && !(await webgpuUsable())) {
        throw new Error('no GPU adapter (hardware acceleration off or unavailable)');
      }
      // The encoder is fp32-only: fp16 silently NaNs on some GPU/driver combos,
      // and a single NaN poisons the whole image (renders black, no error).
      const encPath = cfg.text_encoder;
      const encBuf = await download(fileUrl(repo, encPath), encPath, onProgress);
      const opts = { executionProviders: [backend] };
      const [dit, enc, vae] = await Promise.all([
        ort.InferenceSession.create(ditBuf, opts),
        ort.InferenceSession.create(encBuf, opts),
        ort.InferenceSession.create(vaeBuf, opts),
      ]);
      selected = { ort, cfg, dit, enc, vae, tokenize, backend, encoderPath: encPath, repo };
      break;
    } catch (err) {
      errors.push(`${backend}: ${err.message}`);
    }
  }
  if (!selected) {
    throw new Error(
      `No available backend found (${errors.join(' | ')}). For WebGPU use Chrome/Edge 113+ ` +
      `with hardware acceleration, Firefox Nightly with dom.webgpu.enabled, or Safari Technology ` +
      `Preview. Pass backends: ['webgpu', 'wasm'] for a slow CPU fallback.`,
    );
  }
  return selected;
}

function statsOf(a) {
  let nan = 0;
  let mx = -Infinity;
  for (let i = 0; i < a.length; i++) {
    const v = a[i];
    if (Number.isNaN(v)) nan++;
    else if (v > mx) mx = v;
  }
  return { n: a.length, nan, max: mx === -Infinity ? null : Math.round(mx * 1000) / 1000 };
}

export async function generate(model, prompt, { seed = 1, steps, cfg: guide, onProgress } = {}) {
  const { ort, cfg, dit, enc, vae, tokenize } = model;
  steps ??= cfg.default_steps;
  guide ??= cfg.default_cfg;
  const N = cfg.latent_ch * cfg.latent_size ** 2;
  const tick = (phase, step = 0) => onProgress?.({ phase, step, steps });

  async function encode(text) {    const t = await tokenize([text], { padding: 'max_length', truncation: true, max_length: cfg.ctx_len, return_attention_mask: true });
    const toBig = (a) => BigInt64Array.from(a.data, (v) => BigInt(v));
    const ids = toBig(t.input_ids);
    const am = toBig(t.attention_mask);
    let maxId = 0n;
    let maskOnes = 0;
    for (let i = 0; i < ids.length; i++) if (ids[i] > maxId) maxId = ids[i];
    for (let i = 0; i < am.length; i++) if (am[i] !== 0n) maskOnes++;
    const tokInfo = `tokens len=${ids.length}/${am.length} maxId=${maxId} maskOnes=${maskOnes}`;
    const out = await enc.run({
      input_ids: new ort.Tensor('int64', ids, [1, cfg.ctx_len]),
      attention_mask: new ort.Tensor('int64', am, [1, cfg.ctx_len]),
    });
    const hidden = Object.values(out)[0];
    // Encoder output may be fp16 halves or fp32 — normalize to float32.
    const raw = hidden.data;
    const data = raw instanceof Uint16Array ? Float32Array.from(raw, halfToFloat) : Float32Array.from(raw);
    const mask = new Float32Array(cfg.ctx_len);
    for (let i = 0; i < cfg.ctx_len; i++) mask[i] = am[i] === 0n ? 0 : 1;
    return { ctx: data, mask, tokInfo };
  }

  // NOTE: the encoder is fp32-only (fp16 silently NaNs on some GPU/driver combos).
  tick('encode');
  const cond = await encode(prompt);
  const uncond = await encode('');
  // Firewall: NaN here would otherwise paint a black image with no error.
  for (const [k, c] of [['cond', cond], ['uncond', uncond]]) {
    for (let i = 0; i < c.ctx.length; i++) {
      if (Number.isNaN(c.ctx[i])) throw new Error(`text encoder returned NaN on this backend (${k} ${i}/${c.ctx.length}) — try another browser`);
    }
  }
  const encDiag = { cond: { ...statsOf(cond.ctx), tok: cond.tokInfo }, uncond: { ...statsOf(uncond.ctx), tok: uncond.tokInfo } };
  const shape = [1, cfg.latent_ch, cfg.latent_size, cfg.latent_size];
  const toCtx = (c) => new ort.Tensor('float32', c.ctx, [1, cfg.ctx_len, c.ctx.length / cfg.ctx_len]);
  let z = gaussian([...prompt].reduce((a, c) => a + c.codePointAt(0), seed), N);
  const dt = 1 / steps;
  let step0Diag = null;
  for (let i = 0; i < steps; i++) {
    const t = new ort.Tensor('float32', new Float32Array([i * dt]), [1]);
    const vs = [];
    for (const c of [cond, uncond]) {
      const out = await dit.run({
        z: new ort.Tensor('float32', z, shape), t,
        ctx: toCtx(c), ctx_mask: new ort.Tensor('float32', c.mask, [1, cfg.ctx_len]),
      });
      vs.push(Object.values(out)[0].data);
    }
    const [vc, vu] = vs;
    const next = new Float32Array(N);
    for (let j = 0; j < N; j++) next[j] = z[j] + dt * (vu[j] + guide * (vc[j] - vu[j]));
    z = next;
    if (i === 0) step0Diag = { vc: statsOf(vc), vu: statsOf(vu) };
    tick('denoise', i + 1);
  }
  tick('decode');
  const scaled = new Float32Array(N);
  for (let i = 0; i < N; i++) scaled[i] = z[i] / cfg.vae_scale;
  const img = await vae.run({ z: new ort.Tensor('float32', scaled, shape) });
  // float NCHW in [-1, 1]; diag locates NaN birth (encoder vs step-0 DiT vs VAE)
  return { pixels: Object.values(img)[0].data, size: cfg.image_size, diag: { enc: encDiag, step0: step0Diag } };
}

export function paint(pixels, size, canvas) {
  canvas.width = size;
  canvas.height = size;
  const ctx = canvas.getContext('2d');
  const img = ctx.createImageData(size, size);
  for (let i = 0; i < size * size; i++) {
    img.data[i * 4] = Math.round(Math.min(1, Math.max(0, (pixels[i] + 1) / 2)) * 255);
    img.data[i * 4 + 1] = Math.round(Math.min(1, Math.max(0, (pixels[size * size + i] + 1) / 2)) * 255);
    img.data[i * 4 + 2] = Math.round(Math.min(1, Math.max(0, (pixels[2 * size * size + i] + 1) / 2)) * 255);
    img.data[i * 4 + 3] = 255;
  }
  ctx.putImageData(img, 0, 0);
}

// Usage:
// const model = await loadSupra(REPO, { onProgress: (p) => console.log(p.file, p.loaded) });
// const { pixels, size } = await generate(model, 'a lighthouse above violet clouds at dusk');
// paint(pixels, size, document.querySelector('canvas'));