Vector Storage: summarize with WebLLM extension

4888e3c2b0af1cec025114c1573a6de6aae65bf6

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

2 files changed, +45 -1Showing whitespace changes
public/scripts/extensions/vectors/index.js+43 -0
@@ -30,6 +30,12 @@ import { textgen_types, textgenerationwebui_settings } from '../../textgen-setti
3030import { SlashCommandParser } from '../../slash-commands/SlashCommandParser.js';
3131import { SlashCommand } from '../../slash-commands/SlashCommand.js';
3232import { ARGUMENT_TYPE, SlashCommandArgument, SlashCommandNamedArgument } from '../../slash-commands/SlashCommandArgument.js';
33+import { generateWebLlmChatPrompt, isWebLlmSupported } from '../shared.js';
34+
35+/**
36+ * @typedef {object} HashedMessage
37+ * @property {string} text - The hashed message text
38+ */
3339
3440const MODULE_NAME = 'vectors';
3541
@@ -191,6 +197,11 @@ function splitByChunks(items) {
191197 return chunkedItems;
192198}
193199
200+/**
201+ * Summarizes messages using the Extras API method.
202+ * @param {HashedMessage[]} hashedMessages Array of hashed messages
203+ * @returns {Promise<HashedMessage[]>} Summarized messages
204+ */
194205async function summarizeExtra(hashedMessages) {
195206 for (const element of hashedMessages) {
196207 try {
@@ -222,6 +233,11 @@ async function summarizeExtra(hashedMessages) {
222233 return hashedMessages;
223234}
224235
236+/**
237+ * Summarizes messages using the main API method.
238+ * @param {HashedMessage[]} hashedMessages Array of hashed messages
239+ * @returns {Promise<HashedMessage[]>} Summarized messages
240+ */
225241async function summarizeMain(hashedMessages) {
226242 for (const element of hashedMessages) {
227243 element.text = await generateRaw(element.text, '', false, false, settings.summary_prompt);
@@ -230,12 +246,39 @@ async function summarizeMain(hashedMessages) {
230246 return hashedMessages;
231247}
232248
249+/**
250+ * Summarizes messages using WebLLM.
251+ * @param {HashedMessage[]} hashedMessages Array of hashed messages
252+ * @returns {Promise<HashedMessage[]>} Summarized messages
253+ */
254+async function summarizeWebLLM(hashedMessages) {
255+ if (!isWebLlmSupported()) {
256+ console.warn('Vectors: WebLLM is not supported');
257+ return hashedMessages;
258+ }
259+
260+ for (const element of hashedMessages) {
261+ const messages = [{ role:'system', content: settings.summary_prompt }, { role:'user', content: element.text }];
262+ element.text = await generateWebLlmChatPrompt(messages);
263+ }
264+
265+ return hashedMessages;
266+}
267+
268+/**
269+ * Summarizes messages using the chosen method.
270+ * @param {HashedMessage[]} hashedMessages Array of hashed messages
271+ * @param {string} endpoint Type of endpoint to use
272+ * @returns {Promise<HashedMessage[]>} Summarized messages
273+ */
233274async function summarize(hashedMessages, endpoint = 'main') {
234275 switch (endpoint) {
235276 case 'main':
236277 return await summarizeMain(hashedMessages);
237278 case 'extras':
238279 return await summarizeExtra(hashedMessages);
280+ case 'webllm':
281+ return await summarizeWebLLM(hashedMessages);
239282 default:
240283 console.error('Unsupported endpoint', endpoint);
241284 }
public/scripts/extensions/vectors/settings.html+2 -1
@@ -378,10 +378,11 @@
378378 <select id="vectors_summary_source" class="text_pole">
379379 <option value="main" data-i18n="Main API">Main API</option>
380380 <option value="extras" data-i18n="Extras API">Extras API</option>
381+ <option value="webllm" data-i18n="WebLLM Extension">WebLLM Extension</option>
381382 </select>
382383
383384 <label for="vectors_summary_prompt" title="Summary Prompt:">Summary Prompt:</label>
384385 <small data-i18n="Only used when Main API or WebLLM Extension is selected.">Only used when Main API or WebLLM Extension is selected.</small>
385386 <textarea id="vectors_summary_prompt" class="text_pole textarea_compact" rows="6" placeholder="This prompt will be sent to AI to request the summary generation."></textarea>
386387 </div>
387388 </div>