branch-chat / src /workers /llmWorker.js
suvadityamuk's picture
suvadityamuk HF Staff
fix: remove ghost text labels in graph view
7334a07
Raw
History Blame Contribute Delete
4.47 kB
/**
* 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);
}
});