File size: 12,982 Bytes
901f944
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
3ec58b3
901f944
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
e21174a
91d3a62
06c4813
74f9ff3
901f944
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
06c4813
 
 
 
 
 
 
 
 
 
 
 
 
 
 
901f944
 
 
 
 
 
 
 
 
 
 
 
 
06c4813
 
 
901f944
 
 
 
 
 
 
f306c40
 
 
 
 
 
 
 
 
 
 
 
901f944
 
 
91d3a62
901f944
 
3ec58b3
901f944
 
 
 
91d3a62
c95d933
 
 
 
 
 
91d3a62
c95d933
f306c40
 
e21174a
3ec58b3
e21174a
3ec58b3
c95d933
 
91d3a62
 
 
 
c95d933
 
91d3a62
c95d933
91d3a62
 
c95d933
 
 
 
 
 
 
e21174a
c95d933
 
 
b28d5d8
 
 
 
 
e21174a
c95d933
 
 
 
 
 
 
74f9ff3
c95d933
74f9ff3
901f944
 
e21174a
74f9ff3
 
 
 
 
25ab319
31b0c92
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
901f944
 
 
25ab319
31b0c92
901f944
 
 
25ab319
901f944
 
 
25ab319
e21174a
 
 
 
74f9ff3
e21174a
74f9ff3
 
901f944
 
 
 
 
 
 
 
 
 
 
 
 
 
 
31b0c92
901f944
 
 
 
 
 
 
 
 
 
 
 
 
 
 
31b0c92
25ab319
901f944
25ab319
901f944
 
 
31b0c92
901f944
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
// TypeScript inference for Bartholomheow/Supra2-IMG-ONNX (browser, WebGPU).
//
//   npm i onnxruntime-web @huggingface/transformers
//   npx tsc --module nodenext --target es2022 --moduleResolution nodenext example_web.ts
//
// Mirrors the tested pipeline 1:1: cached download -> T5 encode -> Euler CFG
// loop (fp16-aware) -> VAE decode -> canvas.

import * as ort from 'onnxruntime-web';
import { AutoTokenizer } from '@huggingface/transformers';

export const REPO = 'Bartholomheow/Supra2-IMG-ONNX';

export interface PipelineConfig {
  dit: string;
  dit_dtype: string;
  text_encoder: string;
  text_encoder_fp32?: string;
  vae_decoder: string;
  latent_size: number;
  latent_ch: number;
  ctx_len: number;
  vae_scale: number;
  image_size: number;
  default_steps: number;
  default_cfg: number;
}

export interface DownloadProgress {
  file: string;
  loaded: number;
  total: number | null;
}

export interface SupraModel {
  cfg: PipelineConfig;
  dit: ort.InferenceSession;
  enc: ort.InferenceSession;
  vae: ort.InferenceSession;
  tokenize: (text: string) => Promise<{ ctx: Float32Array; mask: Float32Array; tokInfo: string }>;
  backend: string;
  encoderPath: string;
  repo: string;
}

const fileUrl = (repo: string, path: string) => `https://huggingface.co/${repo}/resolve/main/${path}`;

function hashSeed(text: string): number {
  let h = 2166136261;
  for (let i = 0; i < text.length; i++) {
    h ^= text.codePointAt(i) ?? 0;
    h = Math.imul(h, 16777619);
  }
  return h >>> 0;
}

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

export function halfToFloat(h: number): number {
  const f = new Float32Array(1);
  const u = new Uint32Array(f.buffer);
  const s = (h & 0x8000) << 16;
  const e = (h >> 10) & 0x1f;
  const m = h & 0x3ff;
  if (e === 0) {
    if (m === 0) { u[0] = s; return f[0]; }
    let mm = m;
    let 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];
}

export function float32ToFloat16(src: Float32Array): Uint16Array {
  const f32 = new Float32Array(1);
  const u32 = new Uint32Array(f32.buffer);
  const out = new Uint16Array(src.length);
  for (let i = 0; i < src.length; i++) {
    f32[0] = src[i];
    const x = u32[0];
    const sign = (x >> 16) & 0x8000;
    const exp = ((x >> 23) & 0xff) - 112;
    out[i] = exp <= 0 ? sign : exp >= 31 ? sign | 0x7bff : sign | (exp << 10) | ((x & 0x7fffff) >> 13);
  }
  return out;
}

async function download(url: string, file: string, onProgress?: (p: DownloadProgress) => void): Promise<ArrayBuffer> {
  const cache = await caches.open('supra2-img-v1');
  const hit = await cache.match(url);
  if (hit) {
    const buf = await hit.arrayBuffer();
    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;
    }
  }
  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: Uint8Array[] = [];
  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;
}

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

export async function loadSupra(
  repo = REPO,
  onProgress?: (p: DownloadProgress) => void,
  backends: string[] = ['webgpu'],
): Promise<SupraModel> {
  const cfg = (await (await fetch(fileUrl(repo, 'pipeline_config.json'))).json()) as PipelineConfig;
  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 tok = await AutoTokenizer.from_pretrained(repo);
  const errors: string[] = [];
  let backend = '';
  let encPath = '';
  let dit!: ort.InferenceSession;
  let enc!: ort.InferenceSession;
  let vae!: ort.InferenceSession;
  for (const b of backends) {
    try {
      if (b === '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).
      encPath = cfg.text_encoder;
      const encBuf = await download(fileUrl(repo, encPath), encPath, onProgress);
      const opts: ort.InferenceSession.SessionOptions = { executionProviders: [b] };
      [dit, enc, vae] = await Promise.all([
        ort.InferenceSession.create(ditBuf, opts),
        ort.InferenceSession.create(encBuf, opts),
        ort.InferenceSession.create(vaeBuf, opts),
      ]);
      backend = b;
      break;
    } catch (err) {
      errors.push(`${b}: ${(err as Error).message}`);
    }
  }
  if (!backend) {
    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.`,
    );
  }
  const tokenize = async (text: string): Promise<{ ctx: Float32Array; mask: Float32Array; tokInfo: string }> => {
    const t = await tok([text], { padding: 'max_length', truncation: true, max_length: cfg.ctx_len, return_attention_mask: true });
    const ids = BigInt64Array.from(t.input_ids.data as ArrayLike<bigint | number>, (v) => BigInt(v));
    const am = BigInt64Array.from(t.attention_mask.data as ArrayLike<bigint | number>, (v) => BigInt(v));
    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];
    const data = hidden.data instanceof Uint16Array ? Float32Array.from(hidden.data, halfToFloat) : Float32Array.from(hidden.data as Float32Array);
    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 };
  };
  return { cfg, dit, enc, vae, tokenize, backend, encoderPath: encPath, repo };
}

export type GeneratePhase = 'encode' | 'denoise' | 'decode';

export function hasNaN(a: ArrayLike<number>): boolean {
  for (let i = 0; i < a.length; i++) if (Number.isNaN(a[i])) return true;
  return false;
}

export interface StageStats {
  n: number;
  nan: number;
  max: number | null;
}

export interface GenerateDiag {
  enc: { cond: StageStats; uncond: StageStats };
  step0: { vc: StageStats; vu: StageStats } | null;
}

export function statsOf(a: ArrayLike<number>): StageStats {
  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: SupraModel,
  prompt: string,
  opts: { seed?: number; steps?: number; cfg?: number; onProgress?: (info: { phase: GeneratePhase; step: number; steps: number }) => void } = {},
): Promise<{ pixels: Float32Array; size: number; diag: GenerateDiag }> {
  const { cfg, dit, vae, tokenize } = model;
  const steps = opts.steps ?? cfg.default_steps;
  const guide = opts.cfg ?? cfg.default_cfg;
  const tick = (phase: GeneratePhase, step = 0): void => opts.onProgress?.({ phase, step, steps });
  const ditFp16 = cfg.dit_dtype !== 'fp32';
  const shape = [1, cfg.latent_ch, cfg.latent_size, cfg.latent_size];
  const n = cfg.latent_ch * cfg.latent_size * cfg.latent_size;
  tick('encode');
  const cond = await tokenize(prompt);
  const uncond = await tokenize('');
  // Firewall: NaN here would otherwise paint a black image with no error.
  // (The encoder is fp32-only precisely because fp16 silently NaNs on some GPUs.)
  for (const [k, c] of [['cond', cond], ['uncond', uncond]] as const) {
    if (hasNaN(c.ctx)) throw new Error(`text encoder returned NaN on this backend (${k} ${c.tokInfo}) — try another browser`);
  }
  const encDiag = { cond: { ...statsOf(cond.ctx), tok: cond.tokInfo }, uncond: { ...statsOf(uncond.ctx), tok: uncond.tokInfo } };
  const toTensor = (data: Float32Array | Uint16Array, dims: number[], fp16: boolean): ort.Tensor =>
    fp16
      ? new ort.Tensor('float16', data instanceof Uint16Array ? data : float32ToFloat16(data), dims)
      : new ort.Tensor('float32', data instanceof Uint16Array ? Float32Array.from(data, halfToFloat) : data, dims);
  const toCtx = (c: { ctx: Float32Array; mask: Float32Array }): ort.Tensor =>
    toTensor(c.ctx, [1, cfg.ctx_len, c.ctx.length / cfg.ctx_len], ditFp16);
  const rand = mulberry32(hashSeed(prompt) + (opts.seed ?? 0));
  let z = 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();
    z[i] = r * Math.cos(a);
    if (i + 1 < n) z[i + 1] = r * Math.sin(a);
  }
  const dt = 1 / steps;
  let step0Diag: GenerateDiag['step0'] = null;
  for (let i = 0; i < steps; i++) {
    const t = new ort.Tensor('float32', new Float32Array([i * dt]), [1]);
    const vs: Float32Array[] = [];
    for (const c of [cond, uncond]) {
      const out = await dit.run({
        z: toTensor(z, shape, ditFp16), t,
        ctx: toCtx(c), ctx_mask: new ort.Tensor('float32', c.mask, [1, cfg.ctx_len]),
      });
      const v = Object.values(out)[0];
      vs.push(v.data instanceof Uint16Array ? Float32Array.from(v.data, halfToFloat) : Float32Array.from(v.data as Float32Array));
    }
    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) });
  return { pixels: Float32Array.from(Object.values(img)[0].data as Float32Array), size: cfg.image_size, diag: { enc: encDiag, step0: step0Diag } };
}

export function paint(pixels: Float32Array, size: number, canvas: HTMLCanvasElement): void {
  canvas.width = size;
  canvas.height = size;
  const ctx = canvas.getContext('2d');
  if (!ctx) throw new Error('2d canvas unavailable');
  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, (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')!);