Mike0021's picture
Enable browser HTTP commands with CORS-aware errors and bounded requests
39371ea verified
Raw History Blame Contribute Delete
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; } };
}