Download app/src/inference.mjs from Mike0021/MiniCPM5-2B-WebGPU-Pi-HTTP: direct link, hf CLI and curl.
- Browser
- Download file 12 kB
-
https://huggingface.co/spaces/Mike0021/MiniCPM5-2B-WebGPU-Pi-HTTP/resolve/main/app/src/inference.mjs
- Command line
-
hf download hf://spaces/Mike0021/MiniCPM5-2B-WebGPU-Pi-HTTP/app/src/inference.mjs
-
curl -L -o inference.mjs https://huggingface.co/spaces/Mike0021/MiniCPM5-2B-WebGPU-Pi-HTTP/resolve/main/app/src/inference.mjs
12 kB
| 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 <think>\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; } }; | |
| } | |