import { AutoTokenizer, AutoModelForCausalLM, TextStreamer, InterruptableStoppingCriteria, env } from '@huggingface/transformers'; import { AssistantMessageEventStream } from '@earendil-works/pi-ai/utils/event-stream'; import { prepareModelCache, MODEL_ID, REVISION } from './download.mjs'; import { fitContext, parseCompletion, splitThinking } from './protocol.mjs'; import { CONTEXT_LIMIT } from './context-usage.mjs'; import { nucleusProcessor } from './sampling.mjs'; import { TokenRateWindow } from './token-rate.mjs'; const MAX_OUTPUT_TOKENS = 2048; export const localModel = { id: MODEL_ID, name: 'MiniCPM5-2B ยท q4f16', api: 'minicpm-webgpu', provider: 'browser', baseUrl: '', reasoning: true, input: ['text'], contextWindow: CONTEXT_LIMIT, maxTokens: MAX_OUTPUT_TOKENS, cost: { input: 0, output: 0, cacheRead: 0, cacheWrite: 0 } }; export function createInference(notify) { let model, tokenizer, device, loadPromise; const stopping = new InterruptableStoppingCriteria(); async function load({ runtimeURL, local = false, origin, limit128 = false, cachedOnly = false }, signal) { if (model) return; if (loadPromise) return loadPromise; loadPromise = (async () => { const adapter = await navigator.gpu?.requestAdapter({ powerPreference: 'high-performance' }); if (!adapter) throw Error('WebGPU is unavailable. Try an up-to-date browser with GPU acceleration enabled.'); if (!adapter.features.has('shader-f16')) throw Error('This model requires WebGPU shader-f16 support on your device.'); if (limit128) { const requestDevice = GPUAdapter.prototype.requestDevice; GPUAdapter.prototype.requestDevice = function (descriptor = {}) { return requestDevice.call(this, { ...descriptor, requiredLimits: { ...descriptor.requiredLimits, maxStorageBufferBindingSize: 128 * 1024 ** 2, maxBufferSize: 256 * 1024 ** 2 } }); }; } device = { vendor: adapter.info.vendor, architecture: adapter.info.architecture, maxStorageBufferBindingSize: adapter.limits.maxStorageBufferBindingSize, features: [...adapter.features] }; notify({ type: 'device', device }); const allowedLocal = local && ['localhost', '127.0.0.1'].includes(new URL(origin).hostname); env.allowRemoteModels = true; env.allowLocalModels = false; env.remoteHost = allowedLocal ? origin + '/' : 'https://huggingface.co/'; env.remotePathTemplate = allowedLocal ? 'models/{model}/' : `{model}/resolve/${REVISION}/`; const name = allowedLocal ? 'minicpm5-webgpu' : MODEL_ID; const baseURL = env.remoteHost + env.remotePathTemplate.replace('{model}', name); env.customCache = await prepareModelCache(baseURL, { signal, cachedOnly, onProgress: p => notify({ type: 'load_progress', ...p }) }); env.useCustomCache = true; env.useBrowserCache = false; env.backends.onnx.wasm.wasmPaths = runtimeURL; env.backends.onnx.wasm.numThreads = 1; signal.throwIfAborted(); notify({ type: 'load_progress', phase: 'compile' }); [tokenizer, model] = await Promise.all([ AutoTokenizer.from_pretrained(name), AutoModelForCausalLM.from_pretrained(name, { device: 'webgpu', dtype: 'q4f16' }), ]); signal.throwIfAborted(); notify({ type: 'load_progress', phase: 'warmup' }); const input = tokenizer('Hello'); await model.generate({ ...input, max_new_tokens: 1, do_sample: false, top_k: 0 }); signal.throwIfAborted(); const gpu = env.backends.onnx.webgpu.device; device.actualLimits = { maxStorageBufferBindingSize: gpu?.limits.maxStorageBufferBindingSize, maxBufferSize: gpu?.limits.maxBufferSize }; gpu?.lost.then(info => { if (info.reason !== 'destroyed') notify({ type: 'fatal', error: 'GPU device was lost. Reload this page to reload the cached model.' }); }); notify({ type: 'loaded', device, cachedBytes: env.customCache.cachedBytes }); })().catch(async error => { await model?.dispose().catch(() => {}); model = undefined; tokenizer = undefined; if (signal.aborted) throw new DOMException('Stopped.', 'AbortError'); throw error; }).finally(() => { loadPromise = undefined; }); return loadPromise; } function streamFn(piModel, context, options = {}) { const stream = new AssistantMessageEventStream(); const message = { role: 'assistant', api: piModel.api, provider: piModel.provider, model: piModel.id, timestamp: Date.now(), content: [], usage: { input: 0, output: 0, cacheRead: 0, cacheWrite: 0, totalTokens: 0, cost: { input: 0, output: 0, cacheRead: 0, cacheWrite: 0, total: 0 } } }; void (async () => { const abort = () => stopping.interrupt(); const tokenRate = new TokenRateWindow(); let rateTimer; let lastUsageUpdate = 0; const publishUsage = (force = false) => { const now = performance.now(); if (!force && now - lastUsageUpdate < 80) return; lastUsageUpdate = now; notify({ type: 'context_usage', inputTokens: message.usage.input, outputTokens: message.usage.output }); }; try { if (!model) throw Error('Load the model first.'); stopping.reset(); options.signal?.throwIfAborted(); options.signal?.addEventListener('abort', abort, { once: true }); notify({ type: 'inference_activity', phase: 'prefill' }); const maxTokens = Math.min(options.maxTokens ?? MAX_OUTPUT_TOKENS, MAX_OUTPUT_TOKENS); const tools = context.tools ?? []; const templateTools = tools.map(t => ({ type: 'function', function: { name: t.name, description: t.description, parameters: t.parameters } })); const { inputs, dropped } = fitContext(context, messages => tokenizer.apply_chat_template(messages, { tools: templateTools, enable_thinking: true, add_generation_prompt: true, return_dict: true, }), CONTEXT_LIMIT - maxTokens); if (dropped) notify({ type: 'context_trim', dropped }); message.usage.input = inputs.input_ids.dims[1]; message.usage.totalTokens = message.usage.input; publishUsage(true); stream.push({ type: 'start', partial: structuredClone(message) }); let raw = '', visible = '', started = false, thought = '', thinkingStarted = false, thinkingEnded = false; let textIndex = 0; const start = performance.now(); let firstTokenMs; const streamer = new TextStreamer(tokenizer, { skip_prompt: true, skip_special_tokens: false, token_callback_function: tokens => { if (!tokens.length) return; const now = performance.now(); firstTokenMs ??= now - start; tokenRate.add(tokens.length, now); if (rateTimer === undefined) { notify({ type: 'inference_activity', phase: 'decode', rate: null, outputTokens: tokenRate.total }); rateTimer = setInterval(() => notify({ type: 'inference_activity', phase: 'decode', rate: tokenRate.rate(performance.now()), outputTokens: tokenRate.total }), 250); } message.usage.output += tokens.length; message.usage.totalTokens = message.usage.input + message.usage.output; publishUsage(); }, callback_function: delta => { raw += delta; // The prompt already ends in \n. Generated text starts // inside reasoning, so looking only for an opening tag is wrong. const parts = splitThinking(raw, { thinkingPrefilled: true }); const nextThought = parts.complete ? parts.thinking : parts.thinking.slice(0, Math.max(0, parts.thinking.length - 8)); if (nextThought.length > thought.length) { if (!thinkingStarted) { message.content = [{ type: 'thinking', thinking: '' }]; stream.push({ type: 'thinking_start', contentIndex: 0, partial: structuredClone(message) }); thinkingStarted = true; textIndex = 1; } message.content[0].thinking = nextThought; stream.push({ type: 'thinking_delta', contentIndex: 0, delta: nextThought.slice(thought.length), partial: structuredClone(message) }); thought = nextThought; } if (!parts.complete) return; if (thinkingStarted && !thinkingEnded) { stream.push({ type: 'thinking_end', contentIndex: 0, content: thought, partial: structuredClone(message) }); thinkingEnded = true; } // Hold possible tags until completion. Tool XML is shown by tool events. const safe = parts.answer.split('<')[0]; if (safe.length <= visible.length) return; if (!started) { message.content.push({ type: 'text', text: '' }); stream.push({ type: 'text_start', contentIndex: textIndex, partial: structuredClone(message) }); started = true; } message.content[textIndex].text = safe; stream.push({ type: 'text_delta', contentIndex: textIndex, delta: safe.slice(visible.length), partial: structuredClone(message) }); visible = safe; }, }); const output = await model.generate({ ...inputs, max_new_tokens: maxTokens, do_sample: true, temperature: 1.0, top_p: 0.95, top_k: 0, repetition_penalty: 1.0, logits_processor: [nucleusProcessor(0.95)], eos_token_id: [1, 130073], streamer, stopping_criteria: stopping }); const ids = output.tolist()[0].slice(inputs.input_ids.dims[1]).map(Number); message.usage.output = ids.length; message.usage.totalTokens = message.usage.input + ids.length; publishUsage(true); options.signal?.throwIfAborted(); raw = tokenizer.decode(ids, { skip_special_tokens: false }); const content = parseCompletion(raw, tools, { thinkingPrefilled: true }); if (started) stream.push({ type: 'text_end', contentIndex: textIndex, content: visible, partial: structuredClone(message) }); message.content = content; const hasTools = content.some(c => c.type === 'toolCall'); // Never execute even a complete prefix of a truncated tool response. const ended = [1, 130073].includes(ids.at(-1)); if (hasTools && !ended) throw Error('The tool response exceeded the output limit. No tool was executed. Try a smaller edit.'); message.stopReason = hasTools ? 'toolUse' : ended ? 'stop' : 'length'; for (const [index, block] of content.entries()) { if (block.type !== 'toolCall') continue; stream.push({ type: 'toolcall_start', contentIndex: index, partial: structuredClone(message) }); stream.push({ type: 'toolcall_end', contentIndex: index, toolCall: block, partial: structuredClone(message) }); } notify({ type: 'generation', raw, thinking: true, sampling: { temperature: 1, topP: 0.95, topK: 0, minP: 0, repetitionPenalty: 1 }, inputTokens: message.usage.input, outputTokens: ids.length, elapsedMs: performance.now() - start, firstTokenMs, stopReason: message.stopReason }); stream.push({ type: 'done', reason: message.stopReason, message }); stream.end(message); } catch (error) { if (message.usage.input) publishUsage(true); message.content = message.content.filter(c => c.type !== 'toolCall'); message.stopReason = options.signal?.aborted ? 'aborted' : 'error'; message.errorMessage = options.signal?.aborted ? 'Stopped.' : String(error.message ?? error); stream.push({ type: 'error', reason: message.stopReason, error: message }); stream.end(message); } finally { clearInterval(rateTimer); notify({ type: 'inference_activity', phase: 'end' }); options.signal?.removeEventListener('abort', abort); } })(); return stream; } return { load, streamFn, stop: () => stopping.interrupt(), get device() { return device; } }; }