Vectors WebLLM (#3631) * Add WebLLM support for vectorization * Load models when WebLLM extension installed * Consistency updated * Move checkWebLlm to initEngine * Refactor vector request handling to use getAdditionalArgs * Add error handling for unsupported WebLLM extension * Add prefix to error causes

1cb9287684fd40821f7e784635dccf9db391503e

Cohee <18619528+Cohee1207@users.noreply.github.com>

Signed
5 files changed, +225 -8Showing whitespace changes
public/scripts/extensions.js+1 -1
@@ -1070,7 +1070,7 @@ export async function installExtension(url, global) {
10701070 toastr.success(t`Extension '${response.display_name}' by ${response.author} (version ${response.version}) has been installed successfully!`, t`Extension installation successful`);
10711071 console.debug(`Extension "${response.display_name}" has been installed successfully at ${response.extensionPath}`);
10721072 await loadExtensionSettings({}, false, false);
10731073 await eventSource.emit(event_types.EXTENSION_SETTINGS_LOADED, response);
10741074}
10751075
10761076/**
public/scripts/extensions/vectors/index.js+133 -7
@@ -19,6 +19,7 @@ import {
1919 modules,
2020 renderExtensionTemplateAsync,
2121 doExtrasFetch, getApiUrl,
22+ openThirdPartyExtensionMenu,
2223} from '../../extensions.js';
2324import { collapseNewlines, registerDebugFunction } from '../../power-user.js';
2425import { SECRET_KEYS, secret_state, writeSecret } from '../../secrets.js';
@@ -34,6 +35,7 @@ import { SlashCommandEnumValue, enumTypes } from '../../slash-commands/SlashComm
3435import { slashCommandReturnHelper } from '../../slash-commands/SlashCommandReturnHelper.js';
3536import { callGenericPopup, POPUP_RESULT, POPUP_TYPE } from '../../popup.js';
3637import { generateWebLlmChatPrompt, isWebLlmSupported } from '../shared.js';
38+import { WebLlmVectorProvider } from './webllm.js';
3739
3840/**
3941 * @typedef {object} HashedMessage
@@ -60,6 +62,7 @@ const settings = {
6062 ollama_model: 'mxbai-embed-large',
6163 ollama_keep: false,
6264 vllm_model: '',
65+ webllm_model: '',
6366 summarize: false,
6467 summarize_sent: false,
6568 summary_source: 'main',
@@ -103,7 +106,7 @@ const settings = {
103106};
104107
105108const moduleWorker = new ModuleWorkerWrapper(synchronizeChat);
106-
109+const webllmProvider = new WebLlmVectorProvider();
107110const cachedSummaries = new Map();
108111
109112/**
@@ -373,6 +376,8 @@ async function synchronizeChat(batchSize = 5) {
373376 return 'Vectorization Source Model is required, but not set.';
374377 case 'extras_module_missing':
375378 return 'Extras API must provide an "embeddings" module.';
379+ case 'webllm_not_supported':
380+ return 'WebLLM extension is not installed or the model is not set.';
376381 default:
377382 return 'Check server console for more details';
378383 }
@@ -747,10 +752,11 @@ async function getQueryText(chat, initiator) {
747752
748753/**
749754 * Gets common body parameters for vector requests.
750755 * @returnsparam {object} args Additional arguments
756+ * @returns {object} Request body
751757 */
752758function getVectorsRequestBody(args = {}) {
753759 const body = Object.assign({}, args);
754760 switch (settings.source) {
755761 case 'extras':
756762 body.extrasUrl = extension_settings.apiUrl;
@@ -777,6 +783,9 @@ function getVectorsRequestBody() {
777783 body.apiUrl = textgenerationwebui_settings.server_urls[textgen_types.VLLM];
778784 body.model = extension_settings.vectors.vllm_model;
779785 break;
786+ case 'webllm':
787+ body.model = extension_settings.vectors.webllm_model;
788+ break;
780789 default:
781790 break;
782791 }
@@ -784,6 +793,21 @@ function getVectorsRequestBody() {
784793}
785794
786795/**
796+ * Gets additional arguments for vector requests.
797+ * @param {string[]} items Items to embed
798+ * @returns {Promise<object>} Additional arguments
799+ */
800+async function getAdditionalArgs(items) {
801+ const args = {};
802+ switch (settings.source) {
803+ case 'webllm':
804+ args.embeddings = await createWebLlmEmbeddings(items);
805+ break;
806+ }
807+ return args;
808+}
809+
810+/**
787811 * Gets the saved hashes for a collection
788812* @param {string} collectionId
789813* @returns {Promise<number[]>} Saved hashes
@@ -816,11 +840,12 @@ async function getSavedHashes(collectionId) {
816840async function insertVectorItems(collectionId, items) {
817841 throwIfSourceInvalid();
818842
843+ const args = await getAdditionalArgs(items.map(x => x.text));
819844 const response = await fetch('/api/vector/insert', {
820845 method: 'POST',
821846 headers: getRequestHeaders(),
822847 body: JSON.stringify({
823848 ...getVectorsRequestBody(args),
824849 collectionId: collectionId,
825850 items: items,
826851 source: settings.source,
@@ -858,6 +883,10 @@ function throwIfSourceInvalid() {
858883 if (settings.source === 'extras' && !modules.includes('embeddings')) {
859884 throw new Error('Vectors: Embeddings module missing', { cause: 'extras_module_missing' });
860885 }
886+
887+ if (settings.source === 'webllm' && (!isWebLlmSupported() || !settings.webllm_model)) {
888+ throw new Error('Vectors: WebLLM is not supported', { cause: 'webllm_not_supported' });
889+ }
861890}
862891
863892/**
@@ -890,11 +919,12 @@ async function deleteVectorItems(collectionId, hashes) {
890919 * @returns {Promise<{ hashes: number[], metadata: object[]}>} - Hashes of the results
891920 */
892921async function queryCollection(collectionId, searchText, topK) {
922+ const args = await getAdditionalArgs([searchText]);
893923 const response = await fetch('/api/vector/query', {
894924 method: 'POST',
895925 headers: getRequestHeaders(),
896926 body: JSON.stringify({
897927 ...getVectorsRequestBody(args),
898928 collectionId: collectionId,
899929 searchText: searchText,
900930 topK: topK,
@@ -919,11 +949,12 @@ async function queryCollection(collectionId, searchText, topK) {
919949 * @returns {Promise<Record<string, { hashes: number[], metadata: object[] }>>} - Results mapped to collection IDs
920950 */
921951async function queryMultipleCollections(collectionIds, searchText, topK, threshold) {
952+ const args = await getAdditionalArgs([searchText]);
922953 const response = await fetch('/api/vector/query-multi', {
923954 method: 'POST',
924955 headers: getRequestHeaders(),
925956 body: JSON.stringify({
926957 ...getVectorsRequestBody(args),
927958 collectionIds: collectionIds,
928959 searchText: searchText,
929960 topK: topK,
@@ -1039,6 +1070,72 @@ function toggleSettings() {
10391070 $('#llamacpp_vectorsModel').toggle(settings.source === 'llamacpp');
10401071 $('#vllm_vectorsModel').toggle(settings.source === 'vllm');
10411072 $('#nomicai_apiKey').toggle(settings.source === 'nomicai');
1073+ $('#webllm_vectorsModel').toggle(settings.source === 'webllm');
1074+ if (settings.source === 'webllm') {
1075+ loadWebLlmModels();
1076+ }
1077+}
1078+
1079+/**
1080+ * Executes a function with WebLLM error handling.
1081+ * @param {function(): Promise<T>} func Function to execute
1082+ * @returns {Promise<T>}
1083+ * @template T
1084+ */
1085+async function executeWithWebLlmErrorHandling(func) {
1086+ try {
1087+ return await func();
1088+ } catch (error) {
1089+ console.log('Vectors: Failed to load WebLLM models', error);
1090+ if (!(error instanceof Error)) {
1091+ return;
1092+ }
1093+ switch (error.cause) {
1094+ case 'webllm-not-available':
1095+ toastr.warning('WebLLM is not available. Please install the extension.', 'WebLLM not installed');
1096+ break;
1097+ case 'webllm-not-updated':
1098+ toastr.warning('The installed extension version does not support embeddings.', 'WebLLM update required');
1099+ break;
1100+ }
1101+ }
1102+}
1103+
1104+/**
1105+ * Loads and displays WebLLM models in the settings.
1106+ * @returns {Promise<void>}
1107+ */
1108+function loadWebLlmModels() {
1109+ return executeWithWebLlmErrorHandling(() => {
1110+ const models = webllmProvider.getModels();
1111+ $('#vectors_webllm_model').empty();
1112+ for (const model of models) {
1113+ $('#vectors_webllm_model').append($('<option>', { value: model.id, text: model.toString() }));
1114+ }
1115+ if (!settings.webllm_model || !models.some(x => x.id === settings.webllm_model)) {
1116+ if (models.length) {
1117+ settings.webllm_model = models[0].id;
1118+ }
1119+ }
1120+ $('#vectors_webllm_model').val(settings.webllm_model);
1121+ return Promise.resolve();
1122+ });
1123+}
1124+
1125+/**
1126+ * Creates WebLLM embeddings for a list of items.
1127+ * @param {string[]} items Items to embed
1128+ * @returns {Promise<Record<string, number[]>>} Calculated embeddings
1129+ */
1130+async function createWebLlmEmbeddings(items) {
1131+ return executeWithWebLlmErrorHandling(async () => {
1132+ const embeddings = await webllmProvider.embedTexts(items, settings.webllm_model);
1133+ const result = /** @type {Record<string, number[]>} */ ({});
1134+ for (let i = 0; i < items.length; i++) {
1135+ result[items[i]] = embeddings[i];
1136+ }
1137+ return result;
1138+ });
10421139}
10431140
10441141async function onPurgeClick() {
@@ -1567,6 +1664,30 @@ jQuery(async () => {
15671664 $('#dialogue_popup_input').val(presetModel);
15681665 });
15691666
1667+ $('#vectors_webllm_install').on('click', (e) => {
1668+ e.preventDefault();
1669+ e.stopPropagation();
1670+
1671+ if (Object.hasOwn(SillyTavern, 'llm')) {
1672+ toastr.info('WebLLM is already installed');
1673+ return;
1674+ }
1675+
1676+ openThirdPartyExtensionMenu('https://github.com/SillyTavern/Extension-WebLLM');
1677+ });
1678+
1679+ $('#vectors_webllm_model').on('input', () => {
1680+ settings.webllm_model = String($('#vectors_webllm_model').val());
1681+ Object.assign(extension_settings.vectors, settings);
1682+ saveSettingsDebounced();
1683+ });
1684+
1685+ $('#vectors_webllm_load').on('click', async () => {
1686+ if (!settings.webllm_model) return;
1687+ await webllmProvider.loadModel(settings.webllm_model);
1688+ toastr.success('WebLLM model loaded');
1689+ });
1690+
15701691 $('#api_key_nomicai').toggleClass('success', !!secret_state[SECRET_KEYS.NOMICAI]);
15711692
15721693 toggleSettings();
@@ -1578,6 +1699,11 @@ jQuery(async () => {
15781699 eventSource.on(event_types.CHAT_DELETED, purgeVectorIndex);
15791700 eventSource.on(event_types.GROUP_CHAT_DELETED, purgeVectorIndex);
15801701 eventSource.on(event_types.FILE_ATTACHMENT_DELETED, purgeFileVectorIndex);
1702+ eventSource.on(event_types.EXTENSION_SETTINGS_LOADED, async (manifest) => {
1703+ if (settings.source === 'webllm' && manifest?.display_name === 'WebLLM') {
1704+ await loadWebLlmModels();
1705+ }
1706+ });
15811707
15821708 SlashCommandParser.addCommandObject(SlashCommand.fromProps({
15831709 name: 'db-ingest',
public/scripts/extensions/vectors/settings.html+16 -0
@@ -21,7 +21,23 @@
2121 <option value="openai">OpenAI</option>
2222 <option value="togetherai">TogetherAI</option>
2323 <option value="vllm">vLLM</option>
24+ <option value="webllm" data-i18n="WebLLM Extension">WebLLM Extension</option>
25+ </select>
26+ </div>
27+ <div class="flex-container flexFlowColumn" id="webllm_vectorsModel">
28+ <label for="vectors_webllm_model" data-i18n="Vectorization Model">
29+ Vectorization Model
30+ </label>
31+ <div class="flex-container">
32+ <select id="vectors_webllm_model" class="text_pole flex1">
2433 </select>
34+ <div id="vectors_webllm_load" class="menu_button menu_button_icon" title="Verify and load the selected model.">
35+ <i class="fa-solid fa-check-to-slot"></i>
36+ </div>
37+ </div>
38+ <div>
39+ Requires the WebLLM extension to be installed. Click <a href="#" id="vectors_webllm_install">here</a> to install.
40+ </div>
2541 </div>
2642 <div class="flex-container flexFlowColumn" id="ollama_vectorsModel">
2743 <label for="vectors_ollama_model" data-i18n="Vectorization Model">
public/scripts/extensions/vectors/webllm.js+64 -0
@@ -0,0 +1,64 @@
1+export class WebLlmVectorProvider {
2+ /** @type {object?} WebLLM engine */
3+ #engine = null;
4+
5+ constructor() {
6+ this.#engine = null;
7+ }
8+
9+ /**
10+ * Check if WebLLM is available and up-to-date
11+ * @throws {Error} If WebLLM is not available or not up-to-date
12+ */
13+ #checkWebLlm() {
14+ if (!Object.hasOwn(SillyTavern, 'llm')) {
15+ throw new Error('WebLLM is not available', { cause: 'webllm-not-available' });
16+ }
17+
18+ if (typeof SillyTavern.llm.generateEmbedding !== 'function') {
19+ throw new Error('WebLLM is not updated', { cause: 'webllm-not-updated' });
20+ }
21+ }
22+
23+ /**
24+ * Initialize the engine with a model.
25+ * @param {string} modelId Model ID to initialize the engine with
26+ * @returns {Promise<void>} Promise that resolves when the engine is initialized
27+ */
28+ #initEngine(modelId) {
29+ this.#checkWebLlm();
30+ if (!this.#engine) {
31+ this.#engine = SillyTavern.llm.getEngine();
32+ }
33+
34+ return this.#engine.loadModel(modelId);
35+ }
36+
37+ /**
38+ * Get available models.
39+ * @returns {{id:string, toString: function(): string}[]} Array of available models
40+ */
41+ getModels() {
42+ this.#checkWebLlm();
43+ return SillyTavern.llm.getEmbeddingModels();
44+ }
45+
46+ /**
47+ * Generate embeddings for a list of texts.
48+ * @param {string[]} texts Array of texts to generate embeddings for
49+ * @param {string} modelId Model to use for generating embeddings
50+ * @returns {Promise<number[][]>} Array of embeddings for each text
51+ */
52+ async embedTexts(texts, modelId) {
53+ await this.#initEngine(modelId);
54+ return this.#engine.generateEmbedding(texts);
55+ }
56+
57+ /**
58+ * Loads a model into the engine.
59+ * @param {string} modelId Model ID to load
60+ */
61+ async loadModel(modelId) {
62+ await this.#initEngine(modelId);
63+ }
64+}
src/endpoints/vectors.js+11 -0
@@ -31,6 +31,7 @@ const SOURCES = [
3131 'ollama',
3232 'llamacpp',
3333 'vllm',
34+ 'webllm',
3435];
3536
3637/**
@@ -64,6 +65,8 @@ async function getVector(source, sourceSettings, text, isQuery, directories) {
6465 return getVllmVector(text, sourceSettings.apiUrl, sourceSettings.model, directories);
6566 case 'ollama':
6667 return getOllamaVector(text, sourceSettings.apiUrl, sourceSettings.model, sourceSettings.keep, directories);
68+ case 'webllm':
69+ return sourceSettings.embeddings[text];
6770 }
6871
6972 throw new Error(`Unknown vector source ${source}`);
@@ -114,6 +117,9 @@ async function getBatchVector(source, sourceSettings, texts, isQuery, directorie
114117 case 'ollama':
115118 results.push(...await getOllamaBatchVector(batch, sourceSettings.apiUrl, sourceSettings.model, sourceSettings.keep, directories));
116119 break;
120+ case 'webllm':
121+ results.push(...texts.map(x => sourceSettings.embeddings[x]));
122+ break;
117123 default:
118124 throw new Error(`Unknown vector source ${source}`);
119125 }
@@ -179,6 +185,11 @@ function getSourceSettings(source, request) {
179185 return {
180186 model: 'nomic-embed-text-v1.5',
181187 };
188+ case 'webllm':
189+ return {
190+ model: String(request.body.model),
191+ embeddings: request.body.embeddings ?? {},
192+ };
182193 default:
183194 return {};
184195 }