pulpie-webgpu / src /lib /webgpu-engine.js
sizzlebop's picture
feat: initial release of Pulpie WebGPU Hugging Face Space
99c7c9f
Raw History Blame Contribute Delete
8.3 kB
/**
* Pulpie WebUI - WebGPU Inference Engine
*
* Manages WebGPU device lifecycle, adapter inspection (including shader-f16 feature check),
* and ONNX Runtime Web session execution for sizzlebop/pulpie-orange-small-onnx (model_fp16.onnx).
*/
import * as ort from 'onnxruntime-web';
import { TOKEN_IDS } from './tokenizer.js';
// Configure ONNX Runtime Web WASM paths using exact runtime version
ort.env.wasm.wasmPaths = `https://cdn.jsdelivr.net/npm/onnxruntime-web@${ort.env.versions.web || '1.29.0'}/dist/`;
ort.env.wasm.numThreads = 1;
/**
* Inspect client WebGPU capabilities and adapter features.
* @returns {Promise<{
* supported: boolean,
* hasShaderF16?: boolean,
* adapterName?: string,
* vendor?: string,
* architecture?: string,
* reason?: string
* }>}
*/
export async function checkWebGpuSupport() {
if (typeof navigator === 'undefined' || !('gpu' in navigator)) {
return {
supported: false,
reason: 'WebGPU is not supported in this browser (navigator.gpu is unavailable).',
};
}
try {
const adapter = await navigator.gpu.requestAdapter({
powerPreference: 'high-performance',
});
if (!adapter) {
return {
supported: false,
reason: 'No compatible WebGPU hardware adapter found on this system.',
};
}
const info = (await adapter.requestAdapterInfo?.()) || {};
const hasShaderF16 = adapter.features.has('shader-f16');
return {
supported: true,
hasShaderF16,
adapterName: info.device || info.description || 'WebGPU Device',
vendor: info.vendor || 'GPU Vendor',
architecture: info.architecture || 'GPU Architecture',
};
} catch (err) {
return {
supported: false,
reason: `WebGPU detection failed: ${err.message}`,
};
}
}
/**
* Initialize an ONNX Runtime InferenceSession with WebGPU or high-speed WASM fallback.
*
* When shader-f16 is missing or WebGPU is unavailable (e.g. non-secure LAN contexts),
* automatically initializes the WASM execution provider to ensure accurate FP16 inference.
*
* @param {ArrayBuffer} modelBuffer - Binary ArrayBuffer of model_fp16.onnx
* @param {function(string): void} [onStatus] - Optional status callback
* @returns {Promise<{ session: ort.InferenceSession, provider: 'webgpu' | 'wasm', hasShaderF16: boolean }>}
*/
export async function createSession(modelBuffer, onStatus = null) {
const gpuCheck = await checkWebGpuSupport();
// If WebGPU is supported, use hardware acceleration (FP32 is universally supported)
if (gpuCheck.supported) {
if (onStatus) onStatus('Initializing WebGPU hardware acceleration...');
try {
const sessionOptions = {
executionProviders: [
{
name: 'webgpu',
deviceType: 'gpu',
powerPreference: 'high-performance',
},
],
graphOptimizationLevel: 'all',
};
const session = await ort.InferenceSession.create(modelBuffer, sessionOptions);
if (onStatus) onStatus('WebGPU session ready');
return { session, provider: 'webgpu', hasShaderF16: Boolean(gpuCheck.hasShaderF16) };
} catch (err) {
console.warn('[WebGPUEngine] WebGPU session creation failed, falling back to WASM:', err);
}
}
// Fallback: WASM engine (when WebGPU is unavailable, e.g. non-secure LAN contexts)
if (onStatus) {
onStatus('WebGPU unavailable (secure context required). Initializing WASM CPU engine...');
}
try {
const session = await ort.InferenceSession.create(modelBuffer, {
executionProviders: ['wasm'],
graphOptimizationLevel: 'all',
});
if (onStatus) onStatus('WASM CPU session ready');
return { session, provider: 'wasm', hasShaderF16: Boolean(gpuCheck.hasShaderF16) };
} catch (err) {
console.error('[WebGPUEngine] Session creation failed:', err);
throw new Error(`Failed to initialize inference session: ${err.message}`);
}
}
/**
* Backwards-compatible wrapper returning InferenceSession directly.
* @param {ArrayBuffer} modelBuffer
* @param {function(string): void} [onStatus]
* @returns {Promise<ort.InferenceSession>}
*/
export async function createWebGpuSession(modelBuffer, onStatus = null) {
const { session } = await createSession(modelBuffer, onStatus);
return session;
}
/**
* Run block classification inference on packed chunks using WebGPU or WASM.
*
* @param {ort.InferenceSession} session - Active ONNX session
* @param {Array<{ chunkIds: number[], blockIndices: number[] }>} chunks - Packed token chunks
* @param {string[]} itemIds - Array of _item_id strings matching each block index
* @param {Object} [options]
* @param {number} [options.sepTokenId=TOKEN_IDS.SEP] - Token ID for <|sep|>
* @param {function({ currentChunk: number, totalChunks: number }): void} [options.onProgress]
* @returns {Promise<{
* labels: Record<string, 'main' | 'other'>,
* rawPredictions: number[],
* inferenceTimeMs: number,
* totalTokens: number
* }>}
*/
export async function runWebGpuInference(
session,
chunks,
itemIds,
{ sepTokenId = TOKEN_IDS.SEP, onProgress = null } = {}
) {
if (!session) throw new Error('Inference session is not initialized');
if (!chunks || chunks.length === 0) {
return {
labels: {},
rawPredictions: [],
inferenceTimeMs: 0,
totalTokens: 0,
};
}
const rawPredictions = new Array(itemIds.length).fill(0);
let totalInferenceTimeMs = 0;
let totalTokens = 0;
for (let cIdx = 0; cIdx < chunks.length; cIdx++) {
const { chunkIds, blockIndices } = chunks[cIdx];
totalTokens += chunkIds.length;
if (onProgress) {
onProgress({ currentChunk: cIdx + 1, totalChunks: chunks.length });
}
// Prepare int64 tensors
const inputIdsBigInt = new BigInt64Array(chunkIds.map((id) => BigInt(id)));
const attentionMaskBigInt = new BigInt64Array(chunkIds.length).fill(1n);
const inputIdsTensor = new ort.Tensor('int64', inputIdsBigInt, [1, chunkIds.length]);
const attentionMaskTensor = new ort.Tensor('int64', attentionMaskBigInt, [1, chunkIds.length]);
const feeds = {
input_ids: inputIdsTensor,
attention_mask: attentionMaskTensor,
};
const startTime = performance.now();
const outputMap = await session.run(feeds);
const latency = performance.now() - startTime;
totalInferenceTimeMs += latency;
const logitsTensor = outputMap.logits;
if (!logitsTensor || !logitsTensor.data) {
throw new Error("Model did not return 'logits' tensor output");
}
const logitsData = logitsTensor.data; // Float32Array [1, seq_len, 2]
// Find indices of <|sep|> tokens in chunkIds
const sepIndices = [];
for (let i = 0; i < chunkIds.length; i++) {
if (chunkIds[i] === sepTokenId) {
sepIndices.push(i);
}
}
// Evaluate argmax at each sep index
for (let i = 0; i < blockIndices.length; i++) {
if (i < sepIndices.length) {
const sepPos = sepIndices[i];
const logitOther = logitsData[sepPos * 2 + 0];
const logitMain = logitsData[sepPos * 2 + 1];
// Handle NaN protection (if GPU shader lacks float16 support)
let predictedClass = 0;
if (!Number.isNaN(logitOther) && !Number.isNaN(logitMain)) {
predictedClass = logitMain > logitOther ? 1 : 0;
} else {
console.warn(`[WebGPUEngine] NaN logit detected at SEP index ${sepPos}.`);
}
const blockIdx = blockIndices[i];
rawPredictions[blockIdx] = predictedClass;
}
}
// Immediately release WebAssembly / WebGPU tensor allocations
try {
if (inputIdsTensor?.dispose) inputIdsTensor.dispose();
if (attentionMaskTensor?.dispose) attentionMaskTensor.dispose();
if (outputMap) {
for (const key in outputMap) {
if (outputMap[key]?.dispose) outputMap[key].dispose();
}
}
} catch {
// Best-effort buffer release
}
}
// Construct label map
const labels = {};
for (let i = 0; i < itemIds.length; i++) {
const id = itemIds[i];
labels[id] = rawPredictions[i] === 1 ? 'main' : 'other';
}
return {
labels,
rawPredictions,
inferenceTimeMs: Number(totalInferenceTimeMs.toFixed(2)),
totalTokens,
};
}