bmanturner's picture
initial release
37e3e2c
Raw
History Blame Contribute Delete
12 kB
"use strict";
(() => {
const MODEL_SHA256 = "8ddbe216c1cd0416a1f6528fc07169a989e5f8f39569ceab1afbc6d9b2e8e839";
const OUTPUTS = ["group_logits", "machine_logits", "exact_logits"];
const HEADS = ["group", "machine", "exact"];
const COUNTS = { group: 1225, machine: 1544, exact: 1499 };
const IMAGE_SIZE = 256;
const MEAN = [0.485, 0.456, 0.406];
const STD = [0.229, 0.224, 0.225];
const fileInput = document.querySelector("#image-file");
const topKInput = document.querySelector("#top-k");
const runButton = document.querySelector("#run");
const preview = document.querySelector("#preview");
const status = document.querySelector("#status");
const errorBox = document.querySelector("#error");
const resultsSection = document.querySelector("#results");
const resultBodies = Object.fromEntries(HEADS.map((head) => [head, document.querySelector(`#${head}-results`)]));
let sessionPromise;
let metadata;
let previewUrl;
if (typeof ort === "undefined") {
status.textContent = "Provider: unavailable";
errorBox.textContent = "ONNX Runtime Web failed to load";
errorBox.style.display = "block";
runButton.disabled = true;
return;
}
ort.env.wasm.wasmPaths = "https://cdn.jsdelivr.net/npm/onnxruntime-web@1.22.0/dist/";
function clearResults() {
for (const body of Object.values(resultBodies)) body.replaceChildren();
resultsSection.style.display = "none";
}
function clearError() {
errorBox.textContent = "";
errorBox.style.display = "none";
}
function showError(error) {
clearResults();
const message = error instanceof Error ? error.message : String(error);
errorBox.textContent = message || "Unknown browser inference error";
errorBox.style.display = "block";
}
function assert(condition, message) {
if (!condition) throw new Error(message);
}
function sameOriginUrl(relativePath) {
const url = new URL(relativePath, window.location.href);
assert(url.origin === window.location.origin, `Refusing cross-origin asset: ${url.href}`);
return url;
}
async function fetchAsset(relativePath, responseType) {
const url = sameOriginUrl(relativePath);
const response = await fetch(url, { credentials: "same-origin" });
assert(response.ok, `Failed to load ${url.pathname}: HTTP ${response.status}`);
assert(new URL(response.url).origin === window.location.origin, `Asset redirected off origin: ${url.pathname}`);
return responseType === "json" ? response.json() : response.arrayBuffer();
}
function arraysEqual(actual, expected) {
return Array.isArray(actual) && actual.length === expected.length && actual.every((value, index) => value === expected[index]);
}
function validateMetadata(value) {
assert(value && typeof value === "object", "Metadata must be an object");
assert(value.format_version === 3, "Unsupported metadata format");
assert(value.label_schema_version === 2, "Unsupported label schema");
assert(value.label_migration?.canonical_vocabulary_sha256 === "7f0de4a0cc94845d5e6f429ca9c6eac81dbef4e7cdfe5f008267ce19d78c1cc1", "Unexpected canonical vocabulary digest");
assert(value.browser_runtime?.batch_size === 1, "Metadata must specify browser batch size 1");
assert(value.browser_runtime?.wasm_fallback_required === true, "Metadata must require WASM fallback");
assert(arraysEqual(value.browser_runtime?.preferred_execution_providers, ["webgpu", "wasm"]), "Unexpected provider order");
assert(value.onnx?.file === "model.fp16.onnx", "Unexpected model filename");
assert(value.onnx?.sha256 === MODEL_SHA256, "Unexpected metadata model digest");
assert(value.onnx?.fixed_batch_size === 1, "Model must use fixed batch size 1");
assert(value.input?.name === "images", "Unexpected input name");
assert(value.input?.dtype === "float32", "Unexpected input type");
assert(arraysEqual(value.input?.shape, [1, 3, IMAGE_SIZE, IMAGE_SIZE]), "Unexpected input shape");
assert(value.input?.resize_short_edge === IMAGE_SIZE && value.input?.center_crop === IMAGE_SIZE, "Unexpected image sizing contract");
assert(value.input?.color === "RGB" && value.input?.interpolation === "bicubic", "Unexpected image preprocessing contract");
assert(arraysEqual(value.input?.mean, MEAN) && arraysEqual(value.input?.std, STD), "Unexpected normalization contract");
assert(arraysEqual(value.outputs, OUTPUTS), "Unexpected output order");
assert(value.precision?.input === "float32" && value.precision?.outputs === "float32", "Unexpected tensor precision contract");
for (const head of HEADS) {
const vocabulary = value.vocabularies?.[head];
assert(Array.isArray(vocabulary) && vocabulary.length === COUNTS[head], `Unexpected ${head} vocabulary`);
assert(vocabulary.every((id) => typeof id === "string"), `Invalid ${head} vocabulary entry`);
}
}
async function sha256Hex(bytes) {
const digest = await crypto.subtle.digest("SHA-256", bytes);
return Array.from(new Uint8Array(digest), (byte) => byte.toString(16).padStart(2, "0")).join("");
}
async function initializeSession() {
status.textContent = "Provider: loading model and metadata…";
const [loadedMetadata, modelBytes] = await Promise.all([
fetchAsset("./onnx-metadata.json", "json"),
fetchAsset("./model.fp16.onnx", "arrayBuffer")
]);
validateMetadata(loadedMetadata);
const digest = await sha256Hex(modelBytes);
assert(digest === MODEL_SHA256, `Model SHA-256 mismatch: expected ${MODEL_SHA256}, got ${digest}`);
let session;
let provider;
try {
if (new URLSearchParams(window.location.search).get("forceWebgpuFailure") === "1") {
throw new Error("WebGPU creation force-failed by query parameter");
}
session = await ort.InferenceSession.create(modelBytes, { executionProviders: ["webgpu"] });
provider = "webgpu";
} catch (webgpuError) {
status.textContent = "Provider: WebGPU unavailable; initializing WASM…";
try {
session = await ort.InferenceSession.create(modelBytes, { executionProviders: ["wasm"] });
provider = "wasm";
} catch (wasmError) {
throw new Error(`Unable to initialize WebGPU or WASM. WebGPU: ${webgpuError.message}; WASM: ${wasmError.message}`);
}
}
metadata = loadedMetadata;
status.textContent = `Provider: ${provider}`;
return session;
}
function getSession() {
if (!sessionPromise) {
sessionPromise = initializeSession().catch((error) => {
sessionPromise = undefined;
metadata = undefined;
status.textContent = "Provider: unavailable";
throw error;
});
}
return sessionPromise;
}
async function preprocessImage(file) {
let bitmap;
try {
bitmap = await createImageBitmap(file);
} catch (error) {
throw new Error(`Unable to decode the selected image: ${error.message}`);
}
try {
assert(bitmap.width > 0 && bitmap.height > 0, "Decoded image has invalid dimensions");
let resizedWidth;
let resizedHeight;
if (bitmap.width < bitmap.height) {
resizedWidth = IMAGE_SIZE;
resizedHeight = Math.trunc(bitmap.height * IMAGE_SIZE / bitmap.width);
} else {
resizedHeight = IMAGE_SIZE;
resizedWidth = Math.trunc(bitmap.width * IMAGE_SIZE / bitmap.height);
}
const cropX = Math.round((resizedWidth - IMAGE_SIZE) / 2);
const cropY = Math.round((resizedHeight - IMAGE_SIZE) / 2);
const canvas = document.createElement("canvas");
canvas.width = IMAGE_SIZE;
canvas.height = IMAGE_SIZE;
const context = canvas.getContext("2d", { willReadFrequently: true });
assert(context, "Canvas 2D is unavailable");
context.imageSmoothingEnabled = true;
context.imageSmoothingQuality = "high";
context.drawImage(bitmap, -cropX, -cropY, resizedWidth, resizedHeight);
const rgba = context.getImageData(0, 0, IMAGE_SIZE, IMAGE_SIZE).data;
const plane = IMAGE_SIZE * IMAGE_SIZE;
const chw = new Float32Array(3 * plane);
for (let pixel = 0; pixel < plane; pixel += 1) {
const rgbaOffset = pixel * 4;
chw[pixel] = (rgba[rgbaOffset] / 255 - MEAN[0]) / STD[0];
chw[plane + pixel] = (rgba[rgbaOffset + 1] / 255 - MEAN[1]) / STD[1];
chw[2 * plane + pixel] = (rgba[rgbaOffset + 2] / 255 - MEAN[2]) / STD[2];
}
return new ort.Tensor("float32", chw, [1, 3, IMAGE_SIZE, IMAGE_SIZE]);
} finally {
bitmap.close();
}
}
function rankOutput(tensor, head, topK) {
assert(tensor && tensor.type === "float32", `${head} output must be float32`);
assert(arraysEqual(tensor.dims, [1, COUNTS[head]]) || arraysEqual(tensor.dims, [COUNTS[head]]), `${head} output has unexpected dimensions`);
assert(tensor.data.length === COUNTS[head], `${head} output has unexpected length`);
let maximum = -Infinity;
for (const value of tensor.data) {
assert(Number.isFinite(value), `${head} output contains a non-finite logit`);
if (value > maximum) maximum = value;
}
const probabilities = new Float64Array(tensor.data.length);
let denominator = 0;
for (let index = 0; index < tensor.data.length; index += 1) {
const probability = Math.exp(tensor.data[index] - maximum);
probabilities[index] = probability;
denominator += probability;
}
assert(Number.isFinite(denominator) && denominator > 0, `${head} softmax failed`);
return Array.from(probabilities, (probability, index) => ({ index, confidence: probability / denominator }))
.sort((left, right) => right.confidence - left.confidence || left.index - right.index)
.slice(0, topK)
.map(({ index, confidence }) => ({ id: metadata.vocabularies[head][index], confidence }));
}
function render(head, rows) {
const fragment = document.createDocumentFragment();
for (const row of rows) {
const tr = document.createElement("tr");
const id = document.createElement("td");
const confidence = document.createElement("td");
id.textContent = row.id;
confidence.textContent = row.confidence.toFixed(6);
tr.append(id, confidence);
fragment.append(tr);
}
resultBodies[head].replaceChildren(fragment);
}
fileInput.addEventListener("change", () => {
clearError();
clearResults();
if (previewUrl) URL.revokeObjectURL(previewUrl);
const file = fileInput.files?.[0];
if (!file) {
preview.removeAttribute("src");
preview.style.display = "none";
previewUrl = undefined;
return;
}
previewUrl = URL.createObjectURL(file);
preview.src = previewUrl;
preview.style.display = "block";
});
runButton.addEventListener("click", async () => {
clearError();
clearResults();
const file = fileInput.files?.[0];
const topK = Number(topKInput.value);
try {
assert(file, "Select an image before running inference");
assert(Number.isInteger(topK) && topK >= 1 && topK <= 20, "Results per head must be an integer between 1 and 20");
runButton.disabled = true;
const [session, tensor] = await Promise.all([getSession(), preprocessImage(file)]);
status.textContent = `${status.textContent}; running inference…`;
const outputs = await session.run({ images: tensor });
for (let index = 0; index < HEADS.length; index += 1) {
const head = HEADS[index];
const outputName = OUTPUTS[index];
assert(Object.prototype.hasOwnProperty.call(outputs, outputName), `Missing output ${outputName}`);
render(head, rankOutput(outputs[outputName], head, topK));
}
resultsSection.style.display = "grid";
status.textContent = status.textContent.replace("; running inference…", "");
} catch (error) {
status.textContent = status.textContent.replace("; running inference…", "");
showError(error);
} finally {
runButton.disabled = false;
}
});
})();