MNIST Classifier
Draw a digit and classify it with ONNX Runtime Web on WebGPU. Render the GPU-resident logits through a non-owning vgpu buffer wrap.
/**
* ORT-free visualization for the MNIST example.
*
* `scripts/render-example-thumbs.mjs` bundles this module for Node, so it must
* never import ONNX Runtime Web, not even dynamically. Session orchestration
* lives in `ort-runtime.ts`.
*/
import type { Buffer, Effect, Gpu, Surface, Target } from 'vgpu';
import { effect as createEffect, frame as runFrame } from 'vgpu';
import type { ThumbnailOptions } from '../../lib/example-renderer';
import { createFixtureDigit, GOLDEN_LOGITS, LOGIT_BYTES } from './fixtures';
import { INPUT_SIZE } from './preprocess';
import visualizeWgsl from './visualize.wgsl';
/** Byte view for `Buffer.write`; narrows TypeScript's ArrayBufferLike generic. */
function asWriteData(view: Float32Array): Uint8Array<ArrayBuffer> {
return new Uint8Array(view.buffer as ArrayBuffer, view.byteOffset, view.byteLength);
}
export interface Visualizer {
/**
* Draws one frame. `logits` may be a non-owning wrap of ORT's output buffer;
* this function only reads it inside the submitted pass and never retains it.
*/
render(gpu: Gpu, output: Surface | Target, logits: Buffer, digit: Buffer, hasResult: boolean): void;
dispose(): void;
}
export function createVisualizer(gpu: Gpu, label = 'mnist-classifier'): Visualizer {
const effect: Effect = createEffect(gpu, visualizeWgsl, { label: `${label}-visualize` });
return {
render(currentGpu, output, logits, digit, hasResult) {
effect.set({
uniforms: {
resolution: output.size,
has_result: hasResult ? 1 : 0,
input_size: INPUT_SIZE,
},
logits,
digit,
});
runFrame(currentGpu, (frame) => frame.pass({ target: output }, (pass) => pass.draw(effect)));
},
dispose() {
// Effects are owned by the gpu's render service; nothing extra to release today.
},
};
}
/** Storage buffer holding the normalized 28x28 input preview. */
export function createDigitBuffer(gpu: Gpu, label = 'mnist-classifier'): Buffer {
return gpu.device.createBuffer({
size: INPUT_SIZE * INPUT_SIZE * 4,
usage: ['storage', 'copy_dst'],
label: `${label}-digit`,
});
}
/** Storage buffer standing in for logits before the first inference. */
export function createIdleLogitsBuffer(gpu: Gpu, label = 'mnist-classifier'): Buffer {
return gpu.device.createBuffer({
size: LOGIT_BYTES,
usage: ['storage', 'copy_dst'],
label: `${label}-idle-logits`,
});
}
export function writeDigit(buffer: Buffer, pixels: Float32Array): void {
buffer.write(asWriteData(pixels));
}
export function writeLogits(buffer: Buffer, logits: Float32Array): void {
buffer.write(asWriteData(logits));
}
/**
* Deterministic thumbnail: uploads the seeded digit and the golden logits
* captured from the committed model, then runs the production shader.
*
* This validates the visualizer only. It proves nothing about ORT interop,
* which requires a real browser; see the example's browser evidence.
*/
export async function renderThumbnail(
gpu: Gpu,
target: Target,
_options: ThumbnailOptions = {},
): Promise<void> {
const visualizer = createVisualizer(gpu, 'mnist-classifier-thumb');
const digit = createDigitBuffer(gpu, 'mnist-classifier-thumb');
const logits = createIdleLogitsBuffer(gpu, 'mnist-classifier-thumb');
try {
writeDigit(digit, createFixtureDigit());
writeLogits(logits, GOLDEN_LOGITS);
visualizer.render(gpu, target, logits, digit, true);
} finally {
await Promise.allSettled([
Promise.resolve().then(() => gpu.gpu.queue.onSubmittedWorkDone()),
Promise.resolve().then(() => gpu.settled()),
]);
visualizer.dispose();
logits.dispose();
digit.dispose();
}
}