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.

Open fullscreen
/**
 * 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();
  }
}