Spaces:
Sleeping
Sleeping
| /** | |
| * LLM Web Worker | |
| * | |
| * Runs Transformers.js text-generation pipeline in a background thread | |
| * so the UI stays responsive during inference. | |
| * | |
| * Communication protocol: | |
| * IN: { type: 'load' } | |
| * IN: { type: 'generate', nodeId, messages } | |
| * OUT: { type: 'load-step', message } | |
| * OUT: { type: 'load-progress', percent } | |
| * OUT: { type: 'load-ready' } | |
| * OUT: { type: 'load-error', error } | |
| * OUT: { type: 'token', nodeId, token } | |
| * OUT: { type: 'done', nodeId, fullText } | |
| * OUT: { type: 'error', nodeId, error } | |
| */ | |
| import { pipeline, TextStreamer } from '@huggingface/transformers'; | |
| const MODEL_ID = 'onnx-community/Qwen2.5-0.5B-Instruct'; | |
| const MODEL_CONFIG = { | |
| dtype: 'q4', | |
| device: 'webgpu', | |
| }; | |
| let generator = null; | |
| function postStep(message) { | |
| self.postMessage({ type: 'load-step', message }); | |
| } | |
| async function loadModel() { | |
| try { | |
| postStep('Checking WebGPU availability...'); | |
| // Check WebGPU support | |
| if (!navigator.gpu) { | |
| throw new Error( | |
| 'WebGPU is not supported in this browser. Please use Chrome 113+ or Edge 113+.' | |
| ); | |
| } | |
| const adapter = await navigator.gpu.requestAdapter(); | |
| if (!adapter) { | |
| throw new Error( | |
| 'No WebGPU adapter found. Your GPU may not be supported.' | |
| ); | |
| } | |
| let gpuName = 'WebGPU-capable GPU'; | |
| try { | |
| // requestAdapterInfo may not exist in all WebGPU implementations | |
| if (adapter.requestAdapterInfo) { | |
| const adapterInfo = await adapter.requestAdapterInfo(); | |
| gpuName = adapterInfo.description || adapterInfo.vendor || gpuName; | |
| } else if (adapter.info) { | |
| gpuName = adapter.info.description || adapter.info.vendor || gpuName; | |
| } | |
| } catch (_) { | |
| // Ignore — GPU name is non-critical | |
| } | |
| postStep(`GPU detected: ${gpuName}`); | |
| postStep(`Initializing model: ${MODEL_ID}`); | |
| postStep('Downloading model weights (q4 quantized, ~350MB)...'); | |
| postStep('This is a one-time download — cached locally after first load.'); | |
| self.postMessage({ type: 'load-progress', percent: 5 }); | |
| generator = await pipeline('text-generation', MODEL_ID, { | |
| ...MODEL_CONFIG, | |
| progress_callback: (progress) => { | |
| if (progress.status === 'downloading' || progress.status === 'progress') { | |
| const percent = progress.progress | |
| ? Math.round(progress.progress) | |
| : 0; | |
| self.postMessage({ type: 'load-progress', percent: Math.min(percent, 95) }); | |
| if (progress.file) { | |
| postStep(`Downloading: ${progress.file} (${percent}%)`); | |
| } | |
| } else if (progress.status === 'loading') { | |
| postStep('Loading model into GPU memory...'); | |
| self.postMessage({ type: 'load-progress', percent: 90 }); | |
| } else if (progress.status === 'ready') { | |
| postStep('Compiling WebGPU shaders...'); | |
| self.postMessage({ type: 'load-progress', percent: 95 }); | |
| } | |
| }, | |
| }); | |
| postStep('Model loaded successfully! Ready for inference.'); | |
| self.postMessage({ type: 'load-progress', percent: 100 }); | |
| self.postMessage({ type: 'load-ready' }); | |
| } catch (err) { | |
| console.error('[LLM Worker] Load error:', err); | |
| self.postMessage({ type: 'load-error', error: err.message }); | |
| } | |
| } | |
| async function generate(nodeId, messages) { | |
| if (!generator) { | |
| self.postMessage({ | |
| type: 'error', | |
| nodeId, | |
| error: 'Model not loaded', | |
| }); | |
| return; | |
| } | |
| try { | |
| let fullText = ''; | |
| const streamer = new TextStreamer(generator.tokenizer, { | |
| skip_prompt: true, | |
| callback_function: (token) => { | |
| fullText += token; | |
| self.postMessage({ type: 'token', nodeId, token }); | |
| }, | |
| }); | |
| await generator(messages, { | |
| max_new_tokens: 512, | |
| do_sample: true, | |
| temperature: 0.7, | |
| top_p: 0.9, | |
| streamer, | |
| }); | |
| self.postMessage({ type: 'done', nodeId, fullText }); | |
| } catch (err) { | |
| console.error('[LLM Worker] Generation error:', err); | |
| self.postMessage({ | |
| type: 'error', | |
| nodeId, | |
| error: err.message, | |
| }); | |
| } | |
| } | |
| // Message handler | |
| self.addEventListener('message', async (event) => { | |
| const { type, nodeId, messages } = event.data; | |
| switch (type) { | |
| case 'load': | |
| await loadModel(); | |
| break; | |
| case 'generate': | |
| await generate(nodeId, messages); | |
| break; | |
| default: | |
| console.warn('[LLM Worker] Unknown message type:', type); | |
| } | |
| }); | |