Spaces:
Running
Running
Download src/lib/webgpu-engine.js from sizzlebop/pulpie-webgpu: direct link, hf CLI and curl.
- Browser
- Download file 8.3 kB
-
https://huggingface.co/spaces/sizzlebop/pulpie-webgpu/resolve/main/src/lib/webgpu-engine.js
- Command line
-
hf download hf://spaces/sizzlebop/pulpie-webgpu/src/lib/webgpu-engine.js
-
curl -L -o webgpu-engine.js https://huggingface.co/spaces/sizzlebop/pulpie-webgpu/resolve/main/src/lib/webgpu-engine.js
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, | |
| }; | |
| } | |