UNPKG

@wllama/wllama

Version:

WebAssembly binding for llama.cpp - Enabling on-browser LLM inference

552 lines (503 loc) 15.8 kB
// Start the main llama.cpp let wllamaMalloc; let wllamaStart; let wllamaAction; let wllamaExit; let wllamaDebug; let Module = null; let isCompat = false; let lastStack = ''; let isAborted = false; let hasMultithread = false; ////////////////////////////////////////////////////////////// // UTILS ////////////////////////////////////////////////////////////// // send message back to main thread const msg = (data, transfer) => postMessage(data, transfer); // Convert CPP log into JS log const cppLogToJSLog = (line) => { const matched = line.match(/@@(DEBUG|INFO|WARN|ERROR)@@(.*)/); return !!matched ? { level: (matched[1] === 'INFO' ? 'debug' : matched[1]).toLowerCase(), text: matched[2], } : { level: 'log', text: line }; }; const getHeapU8 = () => { const buffer = Module.wasmMemory.buffer; return new Uint8Array(buffer); }; const toSizeT = (num) => { return isCompat ? Number(num) : BigInt(num); }; // Get module config that forwards stdout/err to main thread const getWModuleConfig = (_argMainScriptBlob) => { var pathConfig = RUN_OPTIONS.pathConfig; var pthreadPoolSize = RUN_OPTIONS.nbThread; var argMainScriptBlob = _argMainScriptBlob; isCompat = RUN_OPTIONS.compat; hasMultithread = pthreadPoolSize > 1; msg({ verb: 'console.debug', args: [ `Multithread enabled: ${hasMultithread}, pthreadPoolSize: ${pthreadPoolSize}`, ], }); if (!pathConfig['wllama.wasm']) { throw new Error('"wllama.wasm" is missing in pathConfig'); } return { noInitialRun: true, print: function (text) { if (arguments.length > 1) text = Array.prototype.slice.call(arguments).join(' '); msg({ verb: 'console.log', args: [text] }); }, printErr: function (text) { if (arguments.length > 1) text = Array.prototype.slice.call(arguments).join(' '); if (text.startsWith('@@STACK@@')) { lastStack = text.slice('@@STACK@@'.length); return; } const logLine = cppLogToJSLog(text); msg({ verb: 'console.' + logLine.level, args: [logLine.text] }); }, locateFile: function (filename, basePath) { const p = pathConfig[filename]; const truncate = (str) => str.length > 128 ? `${str.substr(0, 128)}...` : str; if (filename.match(/wllama\.worker\.js/)) { msg({ verb: 'console.error', args: [ '"wllama.worker.js" is removed from v2.2.1. Hint: make sure to clear browser\'s cache.', ], }); } else { msg({ verb: 'console.debug', args: [`Loading "${filename}" from "${truncate(p)}"`], }); return p; } }, mainScriptUrlOrBlob: hasMultithread ? argMainScriptBlob : 'throw new Error("Multithreading is not enabled")', pthreadPoolSize: hasMultithread ? pthreadPoolSize : 0, wasmMemory: hasMultithread ? getWasmMemory() : null, onAbort: function (message) { isAborted = true; msg({ verb: 'signal.abort', args: ['abort', message, lastStack, null] }); }, onExit: function (code) { isAborted = true; const callstack = new Error().stack.toString(); msg({ verb: 'signal.abort', args: ['abort', 'exit(' + code + ')', callstack, null], }); }, }; }; // Get the memory to be used by wasm. (Only used in multi-thread mode) // Because we have a weird OOM issue on iOS, we need to try some values // See: https://github.com/emscripten-core/emscripten/issues/19144 // https://github.com/godotengine/godot/issues/70621 const getWasmMemory = () => { let minBytes = 128 * 1024 * 1024; let maxBytes = 4096 * 1024 * 1024; let stepBytes = 128 * 1024 * 1024; while (maxBytes > minBytes) { try { const wasmMemory = new WebAssembly.Memory({ initial: toSizeT(minBytes / 65536), maximum: toSizeT(maxBytes / 65536), shared: true, address: isCompat ? undefined : 'i64', }); return wasmMemory; } catch (e) { maxBytes -= stepBytes; continue; // retry } } throw new Error('Cannot allocate WebAssembly.Memory'); }; ////////////////////////////////////////////////////////////// // HEAPFS PATCH ////////////////////////////////////////////////////////////// /** * By default, emscripten uses memfs. The way it works is by * allocating new Uint8Array in javascript heap. This is not good * because it requires files to be copied to wasm heap each time * a file is read. * * HeapFS is an alternative, which resolves this problem by * allocating space for file directly inside wasm heap. This * allows us to mmap without doing any copy. * * For llama.cpp, this is great because we use MAP_SHARED * * Ref: https://github.com/ngxson/wllama/pull/39 * Ref: https://github.com/emscripten-core/emscripten/blob/main/src/library_memfs.js * * Note 29/05/2024 @ngxson * Due to ftell() being limited to MAX_LONG, we cannot load files bigger than 2^31 bytes (or 2GB) * Ref: https://github.com/emscripten-core/emscripten/blob/main/system/lib/libc/musl/src/stdio/ftell.c */ const fsNameToFile = {}; // map Name => File const fsIdToFile = {}; // map ID => File let currFileId = 0; // Patch and redirect memfs calls to wllama const patchHeapFS = () => { const m = Module; // save functions m.MEMFS.stream_ops._read = m.MEMFS.stream_ops.read; m.MEMFS.stream_ops._write = m.MEMFS.stream_ops.write; m.MEMFS.stream_ops._llseek = m.MEMFS.stream_ops.llseek; m.MEMFS.stream_ops._allocate = m.MEMFS.stream_ops.allocate; m.MEMFS.stream_ops._mmap = m.MEMFS.stream_ops.mmap; m.MEMFS.stream_ops._msync = m.MEMFS.stream_ops.msync; const patchStream = (stream) => { const name = stream.node.name; if (fsNameToFile[name]) { const f = fsNameToFile[name]; const ptr = Number(f.ptr); stream.node.contents = getHeapU8().subarray(ptr, ptr + f.size); stream.node.usedBytes = f.size; } }; // replace "read" functions m.MEMFS.stream_ops.read = function ( stream, buffer, offset, length, position ) { patchStream(stream); return m.MEMFS.stream_ops._read(stream, buffer, offset, length, position); }; m.MEMFS.ops_table.file.stream.read = m.MEMFS.stream_ops.read; // replace "llseek" functions m.MEMFS.stream_ops.llseek = function (stream, offset, whence) { patchStream(stream); return m.MEMFS.stream_ops._llseek(stream, offset, whence); }; m.MEMFS.ops_table.file.stream.llseek = m.MEMFS.stream_ops.llseek; // replace "mmap" functions m.MEMFS.stream_ops.mmap = function (stream, length, position, prot, flags) { patchStream(stream); const name = stream.node.name; if (fsNameToFile[name]) { const f = fsNameToFile[name]; const mmapPtr = f.ptr + toSizeT(position); return { ptr: mmapPtr, allocated: false, }; } else { return m.MEMFS.stream_ops._mmap(stream, length, position, prot, flags); } }; m.MEMFS.ops_table.file.stream.mmap = m.MEMFS.stream_ops.mmap; // mount FS m.FS.mkdir('/models'); m.FS.mount(m.MEMFS, { root: '.' }, '/models'); }; // Allocate a new file in wllama heapfs, returns file ID const heapfsAlloc = (name, size, allocBuffer) => { if (size < 1) { throw new Error('File size must be bigger than 0'); } const m = Module; const ptr = toSizeT(allocBuffer ? m.mmapAlloc(size) : 0); const file = { ptr: ptr, size: size, id: currFileId++, }; fsIdToFile[file.id] = file; fsNameToFile[name] = file; return file.id; }; // Add new file to wllama heapfs, return number of written bytes const heapfsWrite = (id, buffer, offset) => { if (fsIdToFile[id]) { const { ptr, size } = fsIdToFile[id]; const afterWriteByte = offset + buffer.byteLength; if (afterWriteByte > size) { throw new Error( `File ID ${id} write out of bound, afterWriteByte = ${afterWriteByte} while size = ${size}` ); } getHeapU8().set(buffer, Number(ptr) + offset); return buffer.byteLength; } else { throw new Error(`File ID ${id} not found in heapfs`); } }; ////////////////////////////////////////////////////////////// // ASYNC FILE READ ////////////////////////////////////////////////////////////// let isAwaitReading = false; let pendingReadPromise = null; let pendingReadResolve = null; let pendingReadReject = null; const _stripModelsPrefix = (path) => path.replace(/^\/?models\//, ''); // Called from EM_ASYNC_JS stub in wllama-fs.h (path is already a JS string) const _wllama_js_file_read = async (path, offset, req_size, out_ptr) => { const name = _stripModelsPrefix(path); pendingReadPromise = new Promise((res, rej) => { pendingReadResolve = res; pendingReadReject = rej; }); isAwaitReading = true; postMessage({ verb: 'fs.read_req', args: [name, offset, req_size] }); let data; try { data = await pendingReadPromise; } finally { isAwaitReading = false; pendingReadResolve = null; pendingReadReject = null; } const bytes = new Uint8Array(data); getHeapU8().set(bytes, out_ptr); return toSizeT(bytes.length); }; ////////////////////////////////////////////////////////////// // MAIN CODE ////////////////////////////////////////////////////////////// const callWrapper = (name, ret, args, isAsync) => { const fn = Module.cwrap( name, ret, args, isAsync ? { async: true } : undefined ); return async (action, req) => { // console.log(`Calling ${name} with action:`, action, 'and req:', req); let result; try { if (args.length === 2) { result = isAsync ? await fn(action, req) : fn(action, req); } else { result = fn(); } } catch (ex) { console.error(ex); throw ex; } return result; }; }; // re-entering the wasm while a call is suspended (JSPI / asyncify) corrupts its state, so only one call runs at a time and the rest wait in the queue let wasmCallBusy = false; const wasmCallQueue = []; const runWasmCall = async (callbackId, fn) => { if (isAborted) { // the wasm is dead, fail fast instead of calling into it msg({ callbackId, err: 'wllama has crashed, please reload the module' }); return; } if (wasmCallBusy) { wasmCallQueue.push({ callbackId, fn }); return; } wasmCallBusy = true; try { await fn(); } finally { wasmCallBusy = false; if (isAborted) { // do not touch the wasm again after it aborted; the main thread already rejected the queued tasks wasmCallQueue.length = 0; } else { const next = wasmCallQueue.shift(); if (next) runWasmCall(next.callbackId, next.fn); } } }; const runAction = async (data) => { const { args, callbackId } = data; const argAction = args[0]; const argEncodedMsg = args[1]; try { const inputPtr = await wllamaMalloc(toSizeT(argEncodedMsg.byteLength), 0); // copy data to wasm heap const inputBuffer = new Uint8Array( getHeapU8().buffer, Number(inputPtr), argEncodedMsg.byteLength ); inputBuffer.set(argEncodedMsg, 0); const outputPtr = await wllamaAction(argAction, inputPtr); // length of output buffer is written at the first 4 bytes of input buffer const outputLen = new Uint32Array( getHeapU8().buffer, Number(inputPtr), 1 )[0]; // copy the output buffer to JS heap const outputBuffer = new Uint8Array(outputLen); const outputSrcView = new Uint8Array( getHeapU8().buffer, Number(outputPtr), outputLen ); outputBuffer.set(outputSrcView, 0); // copy it msg({ callbackId, result: outputBuffer }, [outputBuffer.buffer]); } catch (err) { handleError(err); } }; function handleError(err) { // If WASM already aborted, onAbort already sent signal.abort; skip to avoid // re-reporting the resulting WebAssembly.RuntimeError as a JS exception. if (isAborted) return; const message = err ? err.message || String(err) : 'Unknown error'; const stack = err ? err.stack || String(err) : ''; msg({ verb: 'signal.abort', args: ['exception', message, stack, err], }); } onmessage = async (e) => { if (!e.data) return; const { verb, args, callbackId } = e.data; // fs.read_res arrives while wasm is JSPI-suspended; resolve the pending promise. if (verb === 'fs.read_res') { if (pendingReadResolve) { pendingReadResolve(args[0]); } return; } // Guard: while awaiting a file read, reject any other incoming task. if (isAwaitReading) { if (callbackId) { msg({ callbackId, err: 'Worker is suspended waiting for file data (JSPI)', }); } return; } if (!callbackId) { msg({ verb: 'console.error', args: ['callbackId is required', e.data] }); return; } if (verb === 'module.init') { const argMainScriptBlob = args[0]; const argUseAsyncFile = args[1]; try { Module = getWModuleConfig(argMainScriptBlob); Module.preRun = () => { if (argUseAsyncFile) { Module.ENV['USE_ASYNC_FILE'] = '1'; } }; Module.onRuntimeInitialized = () => { // async call once module is ready // init FS patchHeapFS(); // init cwrap const pointer = isCompat ? 'number' : 'bigint'; // TODO: note sure why emscripten cannot bind if there is only 1 argument wllamaMalloc = callWrapper('wllama_malloc', pointer, [ 'number', pointer, ]); wllamaStart = callWrapper('wllama_start', 'string', [], true); wllamaAction = callWrapper( 'wllama_action', pointer, ['string', pointer], true ); wllamaExit = callWrapper('wllama_exit', 'string', []); wllamaDebug = callWrapper('wllama_debug', 'string', []); msg({ callbackId, result: null }); }; wModuleInit(); } catch (err) { handleError(err); } return; } if (verb === 'fs.alloc') { const argFilename = args[0]; const argSize = args[1]; const argAllocBuffer = args[2]; try { // create blank file const emptyBuffer = new ArrayBuffer(0); Module['FS_createDataFile']( '/models', argFilename, emptyBuffer, true, true, true ); // alloc data on heap const fileId = heapfsAlloc(argFilename, argSize, argAllocBuffer); msg({ callbackId, result: { fileId } }); } catch (err) { handleError(err); } return; } if (verb === 'fs.write') { const argFileId = args[0]; const argBuffer = args[1]; const argOffset = args[2]; try { const writtenBytes = heapfsWrite(argFileId, argBuffer, argOffset); msg({ callbackId, result: { writtenBytes } }); } catch (err) { handleError(err); } return; } if (verb === 'wllama.start') { await runWasmCall(callbackId, async () => { try { const result = await wllamaStart(); msg({ callbackId, result }); } catch (err) { handleError(err); } }); return; } if (verb === 'wllama.action') { await runWasmCall(callbackId, () => runAction(e.data)); return; } if (verb === 'wllama.exit') { await runWasmCall(callbackId, async () => { try { const result = await wllamaExit(); msg({ callbackId, result }); } catch (err) { handleError(err); } }); return; } if (verb === 'wllama.debug') { await runWasmCall(callbackId, async () => { try { const result = await wllamaDebug(); msg({ callbackId, result }); } catch (err) { handleError(err); } }); return; } };