/** * 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} */ 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, * 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, }; }