Vector Storage: summarize with WebLLM extension
| @@ -30,6 +30,12 @@ import { textgen_types, textgenerationwebui_settings } from '../../textgen-setti | |||
| 30 | import { SlashCommandParser } from '../../slash-commands/SlashCommandParser.js'; | 30 | import { SlashCommandParser } from '../../slash-commands/SlashCommandParser.js'; |
| 31 | import { SlashCommand } from '../../slash-commands/SlashCommand.js'; | 31 | import { SlashCommand } from '../../slash-commands/SlashCommand.js'; |
| 32 | import { ARGUMENT_TYPE, SlashCommandArgument, SlashCommandNamedArgument } from '../../slash-commands/SlashCommandArgument.js'; | 32 | import { 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 | */ | ||
| 33 | 39 | ||
| 34 | const MODULE_NAME = 'vectors'; | 40 | const MODULE_NAME = 'vectors'; |
| 35 | 41 | ||
| @@ -191,6 +197,11 @@ function splitByChunks(items) { | |||
| 191 | return chunkedItems; | 197 | return chunkedItems; |
| 192 | } | 198 | } |
| 193 | 199 | ||
| 200 | /** | ||
| 201 | * Summarizes messages using the Extras API method. | ||
| 202 | * @param {HashedMessage[]} hashedMessages Array of hashed messages | ||
| 203 | * @returns {Promise<HashedMessage[]>} Summarized messages | ||
| 204 | */ | ||
| 194 | async function summarizeExtra(hashedMessages) { | 205 | async function summarizeExtra(hashedMessages) { |
| 195 | for (const element of hashedMessages) { | 206 | for (const element of hashedMessages) { |
| 196 | try { | 207 | try { |
| @@ -222,6 +233,11 @@ async function summarizeExtra(hashedMessages) { | |||
| 222 | return hashedMessages; | 233 | return hashedMessages; |
| 223 | } | 234 | } |
| 224 | 235 | ||
| 236 | /** | ||
| 237 | * Summarizes messages using the main API method. | ||
| 238 | * @param {HashedMessage[]} hashedMessages Array of hashed messages | ||
| 239 | * @returns {Promise<HashedMessage[]>} Summarized messages | ||
| 240 | */ | ||
| 225 | async function summarizeMain(hashedMessages) { | 241 | async function summarizeMain(hashedMessages) { |
| 226 | for (const element of hashedMessages) { | 242 | for (const element of hashedMessages) { |
| 227 | element.text = await generateRaw(element.text, '', false, false, settings.summary_prompt); | 243 | element.text = await generateRaw(element.text, '', false, false, settings.summary_prompt); |
| @@ -230,12 +246,39 @@ async function summarizeMain(hashedMessages) { | |||
| 230 | return hashedMessages; | 246 | return hashedMessages; |
| 231 | } | 247 | } |
| 232 | 248 | ||
| 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 | */ | ||
| 233 | async function summarize(hashedMessages, endpoint = 'main') { | 274 | async function summarize(hashedMessages, endpoint = 'main') { |
| 234 | switch (endpoint) { | 275 | switch (endpoint) { |
| 235 | case 'main': | 276 | case 'main': |
| 236 | return await summarizeMain(hashedMessages); | 277 | return await summarizeMain(hashedMessages); |
| 237 | case 'extras': | 278 | case 'extras': |
| 238 | return await summarizeExtra(hashedMessages); | 279 | return await summarizeExtra(hashedMessages); |
| 280 | case 'webllm': | ||
| 281 | return await summarizeWebLLM(hashedMessages); | ||
| 239 | default: | 282 | default: |
| 240 | console.error('Unsupported endpoint', endpoint); | 283 | console.error('Unsupported endpoint', endpoint); |
| 241 | } | 284 | } |
| @@ -378,10 +378,11 @@ | |||
| 378 | <select id="vectors_summary_source" class="text_pole"> | 378 | <select id="vectors_summary_source" class="text_pole"> |
| 379 | <option value="main" data-i18n="Main API">Main API</option> | 379 | <option value="main" data-i18n="Main API">Main API</option> |
| 380 | <option value="extras" data-i18n="Extras API">Extras API</option> | 380 | <option value="extras" data-i18n="Extras API">Extras API</option> |
| 381 | <option value="webllm" data-i18n="WebLLM Extension">WebLLM Extension</option> | ||
| 381 | </select> | 382 | </select> |
| 382 | 383 | ||
| 383 | <label for="vectors_summary_prompt" title="Summary Prompt:">Summary Prompt:</label> | 384 | <label for="vectors_summary_prompt" title="Summary Prompt:">Summary Prompt:</label> |
| 384 | <small data-i18n="Only used when Main API is selected.">Only used when Main API is selected.</small> | 385 | <small data-i18n="Only used when Main API or WebLLM Extension is selected.">Only used when Main API or WebLLM Extension is selected.</small> |
| 385 | <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> | 386 | <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> |
| 386 | </div> | 387 | </div> |
| 387 | </div> | 388 | </div> |