Supra2-IMG-ONNX / example_web.ts
Bartholomheow's picture
Upload example_web.ts with huggingface_hub
e21174a verified
Raw History Blame Contribute Delete
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')!);