Add WebLLM extension summarization
| @@ -25,6 +25,7 @@ import { SlashCommandParser } from '../../slash-commands/SlashCommandParser.js'; | |||
| 25 | import { SlashCommand } from '../../slash-commands/SlashCommand.js'; | 25 | import { SlashCommand } from '../../slash-commands/SlashCommand.js'; |
| 26 | import { ARGUMENT_TYPE, SlashCommandArgument, SlashCommandNamedArgument } from '../../slash-commands/SlashCommandArgument.js'; | 26 | import { ARGUMENT_TYPE, SlashCommandArgument, SlashCommandNamedArgument } from '../../slash-commands/SlashCommandArgument.js'; |
| 27 | import { MacrosParser } from '../../macros.js'; | 27 | import { MacrosParser } from '../../macros.js'; |
| 28 | import { countWebLlmTokens, generateWebLlmChatPrompt, getWebLlmContextSize, isWebLlmSupported } from '../shared.js'; | ||
| 28 | export { MODULE_NAME }; | 29 | export { MODULE_NAME }; |
| 29 | 30 | ||
| 30 | const MODULE_NAME = '1_memory'; | 31 | const MODULE_NAME = '1_memory'; |
| @@ -36,6 +37,40 @@ let lastMessageHash = null; | |||
| 36 | let lastMessageId = null; | 37 | let lastMessageId = null; |
| 37 | let inApiCall = false; | 38 | let inApiCall = false; |
| 38 | 39 | ||
| 40 | /** | ||
| 41 | * Count the number of tokens in the provided text. | ||
| 42 | * @param {string} text Text to count tokens for | ||
| 43 | * @returns {Promise<number>} Number of tokens in the text | ||
| 44 | */ | ||
| 45 | async function countSourceTokens(text, padding = 0) { | ||
| 46 | if (extension_settings.memory.source === summary_sources.webllm) { | ||
| 47 | const count = await countWebLlmTokens(text); | ||
| 48 | return count + padding; | ||
| 49 | } | ||
| 50 | |||
| 51 | if (extension_settings.memory.source === summary_sources.extras) { | ||
| 52 | const count = getTextTokens(tokenizers.GPT2, text).length; | ||
| 53 | return count + padding; | ||
| 54 | } | ||
| 55 | |||
| 56 | return await getTokenCountAsync(text, padding); | ||
| 57 | } | ||
| 58 | |||
| 59 | async function getSourceContextSize() { | ||
| 60 | const overrideLength = extension_settings.memory.overrideResponseLength; | ||
| 61 | |||
| 62 | if (extension_settings.memory.source === summary_sources.webllm) { | ||
| 63 | const maxContext = await getWebLlmContextSize(); | ||
| 64 | return overrideLength > 0 ? (maxContext - overrideLength) : Math.round(maxContext * 0.75); | ||
| 65 | } | ||
| 66 | |||
| 67 | if (extension_settings.source === summary_sources.extras) { | ||
| 68 | return 1024; | ||
| 69 | } | ||
| 70 | |||
| 71 | return getMaxContextSize(overrideLength); | ||
| 72 | } | ||
| 73 | |||
| 39 | const formatMemoryValue = function (value) { | 74 | const formatMemoryValue = function (value) { |
| 40 | if (!value) { | 75 | if (!value) { |
| 41 | return ''; | 76 | return ''; |
| @@ -55,6 +90,7 @@ const saveChatDebounced = debounce(() => getContext().saveChat(), debounce_timeo | |||
| 55 | const summary_sources = { | 90 | const summary_sources = { |
| 56 | 'extras': 'extras', | 91 | 'extras': 'extras', |
| 57 | 'main': 'main', | 92 | 'main': 'main', |
| 93 | 'webllm': 'webllm', | ||
| 58 | }; | 94 | }; |
| 59 | 95 | ||
| 60 | const prompt_builders = { | 96 | const prompt_builders = { |
| @@ -130,12 +166,12 @@ function loadSettings() { | |||
| 130 | 166 | ||
| 131 | async function onPromptForceWordsAutoClick() { | 167 | async function onPromptForceWordsAutoClick() { |
| 132 | const context = getContext(); | 168 | const context = getContext(); |
| 133 | const maxPromptLength = getMaxContextSize(extension_settings.memory.overrideResponseLength); | 169 | const maxPromptLength = await getSourceContextSize(); |
| 134 | const chat = context.chat; | 170 | const chat = context.chat; |
| 135 | const allMessages = chat.filter(m => !m.is_system && m.mes).map(m => m.mes); | 171 | const allMessages = chat.filter(m => !m.is_system && m.mes).map(m => m.mes); |
| 136 | const messagesWordCount = allMessages.map(m => extractAllWords(m)).flat().length; | 172 | const messagesWordCount = allMessages.map(m => extractAllWords(m)).flat().length; |
| 137 | const averageMessageWordCount = messagesWordCount / allMessages.length; | 173 | const averageMessageWordCount = messagesWordCount / allMessages.length; |
| 138 | const tokensPerWord = await getTokenCountAsync(allMessages.join('\n')) / messagesWordCount; | 174 | const tokensPerWord = await countSourceTokens(allMessages.join('\n')) / messagesWordCount; |
| 139 | const wordsPerToken = 1 / tokensPerWord; | 175 | const wordsPerToken = 1 / tokensPerWord; |
| 140 | const maxPromptLengthWords = Math.round(maxPromptLength * wordsPerToken); | 176 | const maxPromptLengthWords = Math.round(maxPromptLength * wordsPerToken); |
| 141 | // How many words should pass so that messages will start be dropped out of context; | 177 | // How many words should pass so that messages will start be dropped out of context; |
| @@ -168,15 +204,15 @@ async function onPromptForceWordsAutoClick() { | |||
| 168 | 204 | ||
| 169 | async function onPromptIntervalAutoClick() { | 205 | async function onPromptIntervalAutoClick() { |
| 170 | const context = getContext(); | 206 | const context = getContext(); |
| 171 | const maxPromptLength = getMaxContextSize(extension_settings.memory.overrideResponseLength); | 207 | const maxPromptLength = await getSourceContextSize(); |
| 172 | const chat = context.chat; | 208 | const chat = context.chat; |
| 173 | const allMessages = chat.filter(m => !m.is_system && m.mes).map(m => m.mes); | 209 | const allMessages = chat.filter(m => !m.is_system && m.mes).map(m => m.mes); |
| 174 | const messagesWordCount = allMessages.map(m => extractAllWords(m)).flat().length; | 210 | const messagesWordCount = allMessages.map(m => extractAllWords(m)).flat().length; |
| 175 | const messagesTokenCount = await getTokenCountAsync(allMessages.join('\n')); | 211 | const messagesTokenCount = await countSourceTokens(allMessages.join('\n')); |
| 176 | const tokensPerWord = messagesTokenCount / messagesWordCount; | 212 | const tokensPerWord = messagesTokenCount / messagesWordCount; |
| 177 | const averageMessageTokenCount = messagesTokenCount / allMessages.length; | 213 | const averageMessageTokenCount = messagesTokenCount / allMessages.length; |
| 178 | const targetSummaryTokens = Math.round(extension_settings.memory.promptWords * tokensPerWord); | 214 | const targetSummaryTokens = Math.round(extension_settings.memory.promptWords * tokensPerWord); |
| 179 | const promptTokens = await getTokenCountAsync(extension_settings.memory.prompt); | 215 | const promptTokens = await countSourceTokens(extension_settings.memory.prompt); |
| 180 | const promptAllowance = maxPromptLength - promptTokens - targetSummaryTokens; | 216 | const promptAllowance = maxPromptLength - promptTokens - targetSummaryTokens; |
| 181 | const maxMessagesPerSummary = extension_settings.memory.maxMessagesPerRequest || 0; | 217 | const maxMessagesPerSummary = extension_settings.memory.maxMessagesPerRequest || 0; |
| 182 | const averageMessagesPerPrompt = Math.floor(promptAllowance / averageMessageTokenCount); | 218 | const averageMessagesPerPrompt = Math.floor(promptAllowance / averageMessageTokenCount); |
| @@ -213,8 +249,8 @@ function onSummarySourceChange(event) { | |||
| 213 | 249 | ||
| 214 | function switchSourceControls(value) { | 250 | function switchSourceControls(value) { |
| 215 | $('#memory_settings [data-summary-source]').each((_, element) => { | 251 | $('#memory_settings [data-summary-source]').each((_, element) => { |
| 216 | const source = $(element).data('summary-source'); | 252 | const source = element.dataset.summarySource.split(',').map(s => s.trim()); |
| 217 | $(element).toggle(source === value); | 253 | $(element).toggle(source.includes(value)); |
| 218 | }); | 254 | }); |
| 219 | } | 255 | } |
| 220 | 256 | ||
| @@ -359,6 +395,12 @@ async function onChatEvent() { | |||
| 359 | } | 395 | } |
| 360 | } | 396 | } |
| 361 | 397 | ||
| 398 | if (extension_settings.memory.source === summary_sources.webllm) { | ||
| 399 | if (!isWebLlmSupported()) { | ||
| 400 | return; | ||
| 401 | } | ||
| 402 | } | ||
| 403 | |||
| 362 | const context = getContext(); | 404 | const context = getContext(); |
| 363 | const chat = context.chat; | 405 | const chat = context.chat; |
| 364 | 406 | ||
| @@ -431,8 +473,12 @@ async function forceSummarizeChat() { | |||
| 431 | return ''; | 473 | return ''; |
| 432 | } | 474 | } |
| 433 | 475 | ||
| 434 | toastr.info('Summarizing chat...', 'Please wait'); | 476 | const toast = toastr.info('Summarizing chat...', 'Please wait', { timeOut: 0, extendedTimeOut: 0 }); |
| 435 | const value = await summarizeChatMain(context, true, skipWIAN); | 477 | const value = extension_settings.memory.source === summary_sources.main |
| 478 | ? await summarizeChatMain(context, true, skipWIAN) | ||
| 479 | : await summarizeChatWebLLM(context, true); | ||
| 480 | |||
| 481 | toastr.clear(toast); | ||
| 436 | 482 | ||
| 437 | if (!value) { | 483 | if (!value) { |
| 438 | toastr.warning('Failed to summarize chat'); | 484 | toastr.warning('Failed to summarize chat'); |
| @@ -484,16 +530,25 @@ async function summarizeChat(context) { | |||
| 484 | case summary_sources.main: | 530 | case summary_sources.main: |
| 485 | await summarizeChatMain(context, false, skipWIAN); | 531 | await summarizeChatMain(context, false, skipWIAN); |
| 486 | break; | 532 | break; |
| 533 | case summary_sources.webllm: | ||
| 534 | await summarizeChatWebLLM(context, false); | ||
| 535 | break; | ||
| 487 | default: | 536 | default: |
| 488 | break; | 537 | break; |
| 489 | } | 538 | } |
| 490 | } | 539 | } |
| 491 | 540 | ||
| 492 | async function summarizeChatMain(context, force, skipWIAN) { | 541 | /** |
| 493 | 542 | * Check if the chat should be summarized based on the current conditions. | |
| 543 | * Return summary prompt if it should be summarized. | ||
| 544 | * @param {any} context ST context | ||
| 545 | * @param {boolean} force Summarize the chat regardless of the conditions | ||
| 546 | * @returns {Promise<string>} Summary prompt or empty string | ||
| 547 | */ | ||
| 548 | async function getSummaryPromptForNow(context, force) { | ||
| 494 | if (extension_settings.memory.promptInterval === 0 && !force) { | 549 | if (extension_settings.memory.promptInterval === 0 && !force) { |
| 495 | console.debug('Prompt interval is set to 0, skipping summarization'); | 550 | console.debug('Prompt interval is set to 0, skipping summarization'); |
| 496 | return; | 551 | return ''; |
| 497 | } | 552 | } |
| 498 | 553 | ||
| 499 | try { | 554 | try { |
| @@ -505,17 +560,17 @@ async function summarizeChatMain(context, force, skipWIAN) { | |||
| 505 | waitUntilCondition(() => is_send_press === false, 30000, 100); | 560 | waitUntilCondition(() => is_send_press === false, 30000, 100); |
| 506 | } catch { | 561 | } catch { |
| 507 | console.debug('Timeout waiting for is_send_press'); | 562 | console.debug('Timeout waiting for is_send_press'); |
| 508 | return; | 563 | return ''; |
| 509 | } | 564 | } |
| 510 | 565 | ||
| 511 | if (!context.chat.length) { | 566 | if (!context.chat.length) { |
| 512 | console.debug('No messages in chat to summarize'); | 567 | console.debug('No messages in chat to summarize'); |
| 513 | return; | 568 | return ''; |
| 514 | } | 569 | } |
| 515 | 570 | ||
| 516 | if (context.chat.length < extension_settings.memory.promptInterval && !force) { | 571 | if (context.chat.length < extension_settings.memory.promptInterval && !force) { |
| 517 | console.debug(`Not enough messages in chat to summarize (chat: ${context.chat.length}, interval: ${extension_settings.memory.promptInterval})`); | 572 | console.debug(`Not enough messages in chat to summarize (chat: ${context.chat.length}, interval: ${extension_settings.memory.promptInterval})`); |
| 518 | return; | 573 | return ''; |
| 519 | } | 574 | } |
| 520 | 575 | ||
| 521 | let messagesSinceLastSummary = 0; | 576 | let messagesSinceLastSummary = 0; |
| @@ -539,7 +594,7 @@ async function summarizeChatMain(context, force, skipWIAN) { | |||
| 539 | 594 | ||
| 540 | if (!conditionSatisfied && !force) { | 595 | if (!conditionSatisfied && !force) { |
| 541 | console.debug(`Summary conditions not satisfied (messages: ${messagesSinceLastSummary}, interval: ${extension_settings.memory.promptInterval}, words: ${wordsSinceLastSummary}, force words: ${extension_settings.memory.promptForceWords})`); | 596 | console.debug(`Summary conditions not satisfied (messages: ${messagesSinceLastSummary}, interval: ${extension_settings.memory.promptInterval}, words: ${wordsSinceLastSummary}, force words: ${extension_settings.memory.promptForceWords})`); |
| 542 | return; | 597 | return ''; |
| 543 | } | 598 | } |
| 544 | 599 | ||
| 545 | console.log('Summarizing chat, messages since last summary: ' + messagesSinceLastSummary, 'words since last summary: ' + wordsSinceLastSummary); | 600 | console.log('Summarizing chat, messages since last summary: ' + messagesSinceLastSummary, 'words since last summary: ' + wordsSinceLastSummary); |
| @@ -547,6 +602,63 @@ async function summarizeChatMain(context, force, skipWIAN) { | |||
| 547 | 602 | ||
| 548 | if (!prompt) { | 603 | if (!prompt) { |
| 549 | console.debug('Summarization prompt is empty. Skipping summarization.'); | 604 | console.debug('Summarization prompt is empty. Skipping summarization.'); |
| 605 | return ''; | ||
| 606 | } | ||
| 607 | |||
| 608 | return prompt; | ||
| 609 | } | ||
| 610 | |||
| 611 | async function summarizeChatWebLLM(context, force) { | ||
| 612 | if (!isWebLlmSupported()) { | ||
| 613 | return; | ||
| 614 | } | ||
| 615 | |||
| 616 | const prompt = await getSummaryPromptForNow(context, force); | ||
| 617 | |||
| 618 | if (!prompt) { | ||
| 619 | return; | ||
| 620 | } | ||
| 621 | |||
| 622 | const { rawPrompt, lastUsedIndex } = await getRawSummaryPrompt(context, prompt); | ||
| 623 | |||
| 624 | if (lastUsedIndex === null || lastUsedIndex === -1) { | ||
| 625 | if (force) { | ||
| 626 | toastr.info('To try again, remove the latest summary.', 'No messages found to summarize'); | ||
| 627 | } | ||
| 628 | |||
| 629 | return null; | ||
| 630 | } | ||
| 631 | |||
| 632 | const messages = [ | ||
| 633 | { role: 'system', content: prompt }, | ||
| 634 | { role: 'user', content: rawPrompt }, | ||
| 635 | ]; | ||
| 636 | |||
| 637 | const params = {}; | ||
| 638 | |||
| 639 | if (extension_settings.memory.overrideResponseLength > 0) { | ||
| 640 | params.max_tokens = extension_settings.memory.overrideResponseLength; | ||
| 641 | } | ||
| 642 | |||
| 643 | const summary = await generateWebLlmChatPrompt(messages, params); | ||
| 644 | const newContext = getContext(); | ||
| 645 | |||
| 646 | // something changed during summarization request | ||
| 647 | if (newContext.groupId !== context.groupId || | ||
| 648 | newContext.chatId !== context.chatId || | ||
| 649 | (!newContext.groupId && (newContext.characterId !== context.characterId))) { | ||
| 650 | console.log('Context changed, summary discarded'); | ||
| 651 | return; | ||
| 652 | } | ||
| 653 | |||
| 654 | setMemoryContext(summary, true, lastUsedIndex); | ||
| 655 | return summary; | ||
| 656 | } | ||
| 657 | |||
| 658 | async function summarizeChatMain(context, force, skipWIAN) { | ||
| 659 | const prompt = await getSummaryPromptForNow(context, force); | ||
| 660 | |||
| 661 | if (!prompt) { | ||
| 550 | return; | 662 | return; |
| 551 | } | 663 | } |
| 552 | 664 | ||
| @@ -634,7 +746,7 @@ async function getRawSummaryPrompt(context, prompt) { | |||
| 634 | chat.pop(); // We always exclude the last message from the buffer | 746 | chat.pop(); // We always exclude the last message from the buffer |
| 635 | const chatBuffer = []; | 747 | const chatBuffer = []; |
| 636 | const PADDING = 64; | 748 | const PADDING = 64; |
| 637 | const PROMPT_SIZE = getMaxContextSize(extension_settings.memory.overrideResponseLength); | 749 | const PROMPT_SIZE = await getSourceContextSize(); |
| 638 | let latestUsedMessage = null; | 750 | let latestUsedMessage = null; |
| 639 | 751 | ||
| 640 | for (let index = latestSummaryIndex + 1; index < chat.length; index++) { | 752 | for (let index = latestSummaryIndex + 1; index < chat.length; index++) { |
| @@ -651,7 +763,7 @@ async function getRawSummaryPrompt(context, prompt) { | |||
| 651 | const entry = `${message.name}:\n${message.mes}`; | 763 | const entry = `${message.name}:\n${message.mes}`; |
| 652 | chatBuffer.push(entry); | 764 | chatBuffer.push(entry); |
| 653 | 765 | ||
| 654 | const tokens = await getTokenCountAsync(getMemoryString(true), PADDING); | 766 | const tokens = await countSourceTokens(getMemoryString(true), PADDING); |
| 655 | 767 | ||
| 656 | if (tokens > PROMPT_SIZE) { | 768 | if (tokens > PROMPT_SIZE) { |
| 657 | chatBuffer.pop(); | 769 | chatBuffer.pop(); |
| @@ -13,6 +13,7 @@ | |||
| 13 | <select id="summary_source"> | 13 | <select id="summary_source"> |
| 14 | <option value="main" data-i18n="ext_sum_main_api">Main API</option> | 14 | <option value="main" data-i18n="ext_sum_main_api">Main API</option> |
| 15 | <option value="extras">Extras API</option> | 15 | <option value="extras">Extras API</option> |
| 16 | <option value="webllm" data-i18n="ext_sum_webllm">WebLLM Extension</option> | ||
| 16 | </select><br> | 17 | </select><br> |
| 17 | 18 | ||
| 18 | <div class="flex-container justifyspacebetween alignitemscenter"> | 19 | <div class="flex-container justifyspacebetween alignitemscenter"> |
| @@ -24,7 +25,7 @@ | |||
| 24 | 25 | ||
| 25 | <textarea id="memory_contents" class="text_pole textarea_compact" rows="6" data-i18n="[placeholder]ext_sum_memory_placeholder" placeholder="Summary will be generated here..."></textarea> | 26 | <textarea id="memory_contents" class="text_pole textarea_compact" rows="6" data-i18n="[placeholder]ext_sum_memory_placeholder" placeholder="Summary will be generated here..."></textarea> |
| 26 | <div class="memory_contents_controls"> | 27 | <div class="memory_contents_controls"> |
| 27 | <div id="memory_force_summarize" data-summary-source="main" class="menu_button menu_button_icon" title="Trigger a summary update right now." data-i18n="[title]ext_sum_force_tip"> | 28 | <div id="memory_force_summarize" data-summary-source="main,webllm" class="menu_button menu_button_icon" title="Trigger a summary update right now." data-i18n="[title]ext_sum_force_tip"> |
| 28 | <i class="fa-solid fa-database"></i> | 29 | <i class="fa-solid fa-database"></i> |
| 29 | <span data-i18n="ext_sum_force_text">Summarize now</span> | 30 | <span data-i18n="ext_sum_force_text">Summarize now</span> |
| 30 | </div> | 31 | </div> |
| @@ -58,7 +59,7 @@ | |||
| 58 | <span data-i18n="ext_sum_prompt_builder_3">Classic, blocking</span> | 59 | <span data-i18n="ext_sum_prompt_builder_3">Classic, blocking</span> |
| 59 | </label> | 60 | </label> |
| 60 | </div> | 61 | </div> |
| 61 | <div data-summary-source="main"> | 62 | <div data-summary-source="main,webllm"> |
| 62 | <label for="memory_prompt" class="title_restorable"> | 63 | <label for="memory_prompt" class="title_restorable"> |
| 63 | <span data-i18n="Summary Prompt">Summary Prompt</span> | 64 | <span data-i18n="Summary Prompt">Summary Prompt</span> |
| 64 | <div id="memory_prompt_restore" data-i18n="[title]ext_sum_restore_default_prompt_tip" title="Restore default prompt" class="right_menu_button"> | 65 | <div id="memory_prompt_restore" data-i18n="[title]ext_sum_restore_default_prompt_tip" title="Restore default prompt" class="right_menu_button"> |
| @@ -74,7 +75,7 @@ | |||
| 74 | </label> | 75 | </label> |
| 75 | <input id="memory_override_response_length" type="range" value="{{defaultSettings.overrideResponseLength}}" min="{{defaultSettings.overrideResponseLengthMin}}" max="{{defaultSettings.overrideResponseLengthMax}}" step="{{defaultSettings.overrideResponseLengthStep}}" /> | 76 | <input id="memory_override_response_length" type="range" value="{{defaultSettings.overrideResponseLength}}" min="{{defaultSettings.overrideResponseLengthMin}}" max="{{defaultSettings.overrideResponseLengthMax}}" step="{{defaultSettings.overrideResponseLengthStep}}" /> |
| 76 | <label for="memory_max_messages_per_request"> | 77 | <label for="memory_max_messages_per_request"> |
| 77 | <span data-i18n="ext_sum_raw_max_msg">[Raw] Max messages per request</span> (<span id="memory_max_messages_per_request_value"></span>) | 78 | <span data-i18n="ext_sum_raw_max_msg">[Raw/WebLLM] Max messages per request</span> (<span id="memory_max_messages_per_request_value"></span>) |
| 78 | <small class="memory_disabled_hint" data-i18n="ext_sum_0_unlimited">0 = unlimited</small> | 79 | <small class="memory_disabled_hint" data-i18n="ext_sum_0_unlimited">0 = unlimited</small> |
| 79 | </label> | 80 | </label> |
| 80 | <input id="memory_max_messages_per_request" type="range" value="{{defaultSettings.maxMessagesPerRequest}}" min="{{defaultSettings.maxMessagesPerRequestMin}}" max="{{defaultSettings.maxMessagesPerRequestMax}}" step="{{defaultSettings.maxMessagesPerRequestStep}}" /> | 81 | <input id="memory_max_messages_per_request" type="range" value="{{defaultSettings.maxMessagesPerRequest}}" min="{{defaultSettings.maxMessagesPerRequestMin}}" max="{{defaultSettings.maxMessagesPerRequestMax}}" step="{{defaultSettings.maxMessagesPerRequestStep}}" /> |
| @@ -183,32 +183,40 @@ function throwIfInvalidModel(useReverseProxy) { | |||
| 183 | */ | 183 | */ |
| 184 | export function isWebLlmSupported() { | 184 | export function isWebLlmSupported() { |
| 185 | if (!('gpu' in navigator)) { | 185 | if (!('gpu' in navigator)) { |
| 186 | toastr.error('Your browser does not support the WebGPU API. Please use a different browser.', 'WebLLM', { | 186 | const warningKey = 'webllm_browser_warning_shown'; |
| 187 | preventDuplicates: true, | 187 | if (!sessionStorage.getItem(warningKey)) { |
| 188 | timeOut: 0, | 188 | toastr.error('Your browser does not support the WebGPU API. Please use a different browser.', 'WebLLM', { |
| 189 | extendedTimeOut: 0, | 189 | preventDuplicates: true, |
| 190 | }); | 190 | timeOut: 0, |
| 191 | extendedTimeOut: 0, | ||
| 192 | }); | ||
| 193 | sessionStorage.setItem(warningKey, '1'); | ||
| 194 | } | ||
| 191 | return false; | 195 | return false; |
| 192 | } | 196 | } |
| 193 | 197 | ||
| 194 | if (!('llm' in SillyTavern)) { | 198 | if (!('llm' in SillyTavern)) { |
| 195 | toastr.error('WebLLM extension is not installed. Click here to install it.', 'WebLLM', { | 199 | const warningKey = 'webllm_extension_warning_shown'; |
| 196 | timeOut: 0, | 200 | if (!sessionStorage.getItem(warningKey)) { |
| 197 | extendedTimeOut: 0, | 201 | toastr.error('WebLLM extension is not installed. Click here to install it.', 'WebLLM', { |
| 198 | preventDuplicates: true, | 202 | timeOut: 0, |
| 199 | onclick: () => { | 203 | extendedTimeOut: 0, |
| 200 | const button = document.getElementById('third_party_extension_button'); | 204 | preventDuplicates: true, |
| 201 | if (button) { | 205 | onclick: () => { |
| 202 | button.click(); | 206 | const button = document.getElementById('third_party_extension_button'); |
| 203 | } | 207 | if (button) { |
| 204 | 208 | button.click(); | |
| 205 | const input = document.querySelector('dialog textarea'); | 209 | } |
| 206 | 210 | ||
| 207 | if (input instanceof HTMLTextAreaElement) { | 211 | const input = document.querySelector('dialog textarea'); |
| 208 | input.value = 'https://github.com/SillyTavern/Extension-WebLLM'; | 212 | |
| 209 | } | 213 | if (input instanceof HTMLTextAreaElement) { |
| 210 | }, | 214 | input.value = 'https://github.com/SillyTavern/Extension-WebLLM'; |
| 211 | }); | 215 | } |
| 216 | }, | ||
| 217 | }); | ||
| 218 | sessionStorage.setItem(warningKey, '1'); | ||
| 219 | } | ||
| 212 | return false; | 220 | return false; |
| 213 | } | 221 | } |
| 214 | 222 | ||
| @@ -218,15 +226,16 @@ export function isWebLlmSupported() { | |||
| 218 | /** | 226 | /** |
| 219 | * Generates text in response to a chat prompt using WebLLM. | 227 | * Generates text in response to a chat prompt using WebLLM. |
| 220 | * @param {any[]} messages Messages to use for generating | 228 | * @param {any[]} messages Messages to use for generating |
| 229 | * @param {object} params Additional parameters | ||
| 221 | * @returns {Promise<string>} Generated response | 230 | * @returns {Promise<string>} Generated response |
| 222 | */ | 231 | */ |
| 223 | export async function generateWebLlmChatPrompt(messages) { | 232 | export async function generateWebLlmChatPrompt(messages, params = {}) { |
| 224 | if (!isWebLlmSupported()) { | 233 | if (!isWebLlmSupported()) { |
| 225 | throw new Error('WebLLM extension is not installed.'); | 234 | throw new Error('WebLLM extension is not installed.'); |
| 226 | } | 235 | } |
| 227 | 236 | ||
| 228 | const engine = SillyTavern.llm; | 237 | const engine = SillyTavern.llm; |
| 229 | const response = await engine.generateChatPrompt(messages); | 238 | const response = await engine.generateChatPrompt(messages, params); |
| 230 | return response; | 239 | return response; |
| 231 | } | 240 | } |
| 232 | 241 | ||