Add WebLLM extension summarization

8685c2f471bda4784ed52600620bd7019216a293

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

3 files changed, +145 -23Showing whitespace changes
public/scripts/extensions/memory/index.js+130 -18
@@ -25,6 +25,7 @@ import { SlashCommandParser } from '../../slash-commands/SlashCommandParser.js';
25import { SlashCommand } from '../../slash-commands/SlashCommand.js';25import { SlashCommand } from '../../slash-commands/SlashCommand.js';
26import { ARGUMENT_TYPE, SlashCommandArgument, SlashCommandNamedArgument } from '../../slash-commands/SlashCommandArgument.js';26import { ARGUMENT_TYPE, SlashCommandArgument, SlashCommandNamedArgument } from '../../slash-commands/SlashCommandArgument.js';
27import { MacrosParser } from '../../macros.js';27import { MacrosParser } from '../../macros.js';
28import { countWebLlmTokens, generateWebLlmChatPrompt, getWebLlmContextSize, isWebLlmSupported } from '../shared.js';
28export { MODULE_NAME };29export { MODULE_NAME };
2930
30const MODULE_NAME = '1_memory';31const MODULE_NAME = '1_memory';
@@ -36,6 +37,40 @@ let lastMessageHash = null;
36let lastMessageId = null;37let lastMessageId = null;
37let inApiCall = false;38let inApiCall = false;
3839
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 */
45async 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
59async 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
39const formatMemoryValue = function (value) {74const formatMemoryValue = function (value) {
40 if (!value) {75 if (!value) {
41 return '';76 return '';
@@ -55,6 +90,7 @@ const saveChatDebounced = debounce(() => getContext().saveChat(), debounce_timeo
55const summary_sources = {90const summary_sources = {
56 'extras': 'extras',91 'extras': 'extras',
57 'main': 'main',92 'main': 'main',
93 'webllm': 'webllm',
58};94};
5995
60const prompt_builders = {96const prompt_builders = {
@@ -130,12 +166,12 @@ function loadSettings() {
130166
131async function onPromptForceWordsAutoClick() {167async 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() {
168204
169async function onPromptIntervalAutoClick() {205async 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) {
213249
214function switchSourceControls(value) {250function 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}
220256
@@ -359,6 +395,12 @@ async function onChatEvent() {
359 }395 }
360 }396 }
361397
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;
364406
@@ -431,8 +473,12 @@ async function forceSummarizeChat() {
431 return '';473 return '';
432 }474 }
433475
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);
436482
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}
491540
492async function summarizeChatMain(context, force, skipWIAN) {541/**
493542 * 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 */
548async 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 }
498553
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 }
510565
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 }
515570
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 }
520575
521 let messagesSinceLastSummary = 0;576 let messagesSinceLastSummary = 0;
@@ -539,7 +594,7 @@ async function summarizeChatMain(context, force, skipWIAN) {
539594
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 }
544599
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) {
547602
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
611async 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
658async function summarizeChatMain(context, force, skipWIAN) {
659 const prompt = await getSummaryPromptForNow(context, force);
660
661 if (!prompt) {
550 return;662 return;
551 }663 }
552664
@@ -634,7 +746,7 @@ async function getRawSummaryPrompt(context, prompt) {
634 chat.pop(); // We always exclude the last message from the buffer746 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;
639751
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);
653765
654 const tokens = await getTokenCountAsync(getMemoryString(true), PADDING);766 const tokens = await countSourceTokens(getMemoryString(true), PADDING);
655767
656 if (tokens > PROMPT_SIZE) {768 if (tokens > PROMPT_SIZE) {
657 chatBuffer.pop();769 chatBuffer.pop();
public/scripts/extensions/memory/settings.html+4 -3
@@ -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>
1718
18 <div class="flex-container justifyspacebetween alignitemscenter">19 <div class="flex-container justifyspacebetween alignitemscenter">
@@ -24,7 +25,7 @@
2425
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}}" />
public/scripts/extensions/shared.js+11 -2
@@ -183,15 +183,21 @@ function throwIfInvalidModel(useReverseProxy) {
183 */183 */
184export function isWebLlmSupported() {184export function isWebLlmSupported() {
185 if (!('gpu' in navigator)) {185 if (!('gpu' in navigator)) {
186 const warningKey = 'webllm_browser_warning_shown';
187 if (!sessionStorage.getItem(warningKey)) {
186 toastr.error('Your browser does not support the WebGPU API. Please use a different browser.', 'WebLLM', {188 toastr.error('Your browser does not support the WebGPU API. Please use a different browser.', 'WebLLM', {
187 preventDuplicates: true,189 preventDuplicates: true,
188 timeOut: 0,190 timeOut: 0,
189 extendedTimeOut: 0,191 extendedTimeOut: 0,
190 });192 });
193 sessionStorage.setItem(warningKey, '1');
194 }
191 return false;195 return false;
192 }196 }
193197
194 if (!('llm' in SillyTavern)) {198 if (!('llm' in SillyTavern)) {
199 const warningKey = 'webllm_extension_warning_shown';
200 if (!sessionStorage.getItem(warningKey)) {
195 toastr.error('WebLLM extension is not installed. Click here to install it.', 'WebLLM', {201 toastr.error('WebLLM extension is not installed. Click here to install it.', 'WebLLM', {
196 timeOut: 0,202 timeOut: 0,
197 extendedTimeOut: 0,203 extendedTimeOut: 0,
@@ -209,6 +215,8 @@ export function isWebLlmSupported() {
209 }215 }
210 },216 },
211 });217 });
218 sessionStorage.setItem(warningKey, '1');
219 }
212 return false;220 return false;
213 }221 }
214222
@@ -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 generating228 * @param {any[]} messages Messages to use for generating
229 * @param {object} params Additional parameters
221 * @returns {Promise<string>} Generated response230 * @returns {Promise<string>} Generated response
222 */231 */
223export async function generateWebLlmChatPrompt(messages) {232export 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 }
227236
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}
232241