Download example_web.ts from Bartholomheow/Supra2-IMG-ONNX: direct link, hf CLI and curl.
- Browser
- Download file 13 kB
-
https://huggingface.co/Bartholomheow/Supra2-IMG-ONNX/resolve/main/example_web.ts
- Command line
-
hf download hf://Bartholomheow/Supra2-IMG-ONNX/example_web.ts
-
curl -L -o example_web.ts https://huggingface.co/Bartholomheow/Supra2-IMG-ONNX/resolve/main/example_web.ts
13 kB
| // 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')!); | |