Merge pull request #2988 from QuantumEntangledAndy/feat/cachedVectorSummaries Add client side cacheing of vector summaries

e01a243ce5ee0abd6479d17a2b6e2f58ebf781f5

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

Signed
1 files changed, +84 -73Ignore whitespace
public/scripts/extensions/vectors/index.js+84 -73
@@ -36,6 +36,7 @@ import { generateWebLlmChatPrompt, isWebLlmSupported } from '../shared.js';
3636/**
3737 * @typedef {object} HashedMessage
3838 * @property {string} text - The hashed message text
39+ * @property {number} hash - The hash used as the vector key
3940 */
4041
4142const MODULE_NAME = 'vectors';
@@ -96,6 +97,8 @@ const settings = {
9697
9798const moduleWorker = new ModuleWorkerWrapper(synchronizeChat);
9899
100+const cachedSummaries = new Map();
101+
99102/**
100103 * Gets the Collection ID for a file embedded in the chat.
101104 * @param {string} fileUrl URL of the file
@@ -118,6 +121,10 @@ async function onVectorizeAllClick() {
118121 return;
119122 }
120123
124+ // Clear all cached summaries to ensure that new ones are created
125+ // upon request of a full vectorise
126+ cachedSummaries.clear();
127+
121128 const batchSize = 5;
122129 const elapsedLog = [];
123130 let finished = false;
@@ -200,70 +207,64 @@ function splitByChunks(items) {
200207
201208/**
202209 * Summarizes messages using the Extras API method.
203210 * @param {HashedMessage[]} hashedMessages Array ofelement hashed messagesmessage
204211 * @returns {Promise<HashedMessage[]boolean>} Summarized messagesSucess
205212 */
206213async function summarizeExtra(hashedMessageselement) {
207- for (const element of hashedMessages) {
214+ try {
208- try {
215+ const url = new URL(getApiUrl());
209- const url = new URL(getApiUrl());
216+ url.pathname = '/api/summarize';
210- url.pathname = '/api/summarize';
211-
212- const apiResult = await doExtrasFetch(url, {
213- method: 'POST',
214- headers: {
215- 'Content-Type': 'application/json',
216- 'Bypass-Tunnel-Reminder': 'bypass',
217- },
218- body: JSON.stringify({
219- text: element.text,
220- params: {},
221- }),
222- });
223217
224- if (apiResult.ok) {
218+ const apiResult = await doExtrasFetch(url, {
225- const data = await apiResult.json();
219+ method: 'POST',
226- element.text = data.summary;
220+ headers: {
227- }
221+ 'Content-Type': 'application/json',
228- }
222+ 'Bypass-Tunnel-Reminder': 'bypass',
229- catch (error) {
223+ },
230- console.log(error);
224+ body: JSON.stringify({
225+ text: element.text,
226+ params: {},
227+ }),
228+ });
229+
230+ if (apiResult.ok) {
231+ const data = await apiResult.json();
232+ element.text = data.summary;
231233 }
232234 }
235+ catch (error) {
236+ console.log(error);
237+ return false;
238+ }
233239
234240 return hashedMessagestrue;
235241}
236242
237243/**
238244 * Summarizes messages using the main API method.
239245 * @param {HashedMessage[]} hashedMessages Array ofelement hashed messagesmessage
240246 * @returns {Promise<HashedMessage[]boolean>} Summarized messagesSucess
241247 */
242248async function summarizeMain(hashedMessageselement) {
243- for (const element of hashedMessages) {
249+ element.text = await generateRaw(element.text, '', false, false, settings.summary_prompt);
244- element.text = await generateRaw(element.text, '', false, false, settings.summary_prompt);
250+ return true;
245- }
246-
247- return hashedMessages;
248251}
249252
250253/**
251254 * Summarizes messages using WebLLM.
252255 * @param {HashedMessage[]} hashedMessages Array ofelement hashed messagesmessage
253256 * @returns {Promise<HashedMessage[]boolean>} Summarized messagesSucess
254257 */
255258async function summarizeWebLLM(hashedMessageselement) {
256259 if (!isWebLlmSupported()) {
257260 console.warn('Vectors: WebLLM is not supported');
258261 return hashedMessagesfalse;
259262 }
260263
261- for (const element of hashedMessages) {
264+ const messages = [{ role: 'system', content: settings.summary_prompt }, { role: 'user', content: element.text }];
262- const messages = [{ role:'system', content: settings.summary_prompt }, { role:'user', content: element.text }];
265+ element.text = await generateWebLlmChatPrompt(messages);
263- element.text = await generateWebLlmChatPrompt(messages);
264- }
265266
266267 return hashedMessagestrue;
267268}
268269
269270/**
@@ -273,16 +274,35 @@ async function summarizeWebLLM(hashedMessages) {
273274 * @returns {Promise<HashedMessage[]>} Summarized messages
274275 */
275276async function summarize(hashedMessages, endpoint = 'main') {
276277 switchfor (endpointconst element of hashedMessages) {
277- case 'main':
278+ const cachedSummary = cachedSummaries.get(element.hash);
278- return await summarizeMain(hashedMessages);
279+ if (!cachedSummary) {
279- case 'extras':
280+ let success = true;
280- return await summarizeExtra(hashedMessages);
281+ switch (endpoint) {
281282 case 'webllmmain':
282283 return success = await summarizeWebLLMsummarizeMain(hashedMessageselement);
283- default:
284+ break;
284- console.error('Unsupported endpoint', endpoint);
285+ case 'extras':
286+ success = await summarizeExtra(element);
287+ break;
288+ case 'webllm':
289+ success = await summarizeWebLLM(element);
290+ break;
291+ default:
292+ console.error('Unsupported endpoint', endpoint);
293+ success = false;
294+ break;
295+ }
296+ if (success) {
297+ cachedSummaries.set(element.hash, element.text);
298+ } else {
299+ break;
300+ }
301+ } else {
302+ element.text = cachedSummary;
303+ }
285304 }
305+ return hashedMessages;
286306}
287307
288308async function synchronizeChat(batchSize = 5) {
@@ -307,16 +327,15 @@ async function synchronizeChat(batchSize = 5) {
307327 return -1;
308328 }
309329
310330 letconst hashedMessages = context.chat.filter(x => !x.is_system).map(x => ({ text: String(substituteParams(x.mes)), hash: getStringHash(substituteParams(x.mes)), index: context.chat.indexOf(x) }));
311331 const hashesInCollection = await getSavedHashes(chatId);
312332
313- if (settings.summarize) {
333+ let newVectorItems = hashedMessages.filter(x => !hashesInCollection.includes(x.hash));
314- hashedMessages = await summarize(hashedMessages, settings.summary_source);
315- }
316-
317- const newVectorItems = hashedMessages.filter(x => !hashesInCollection.includes(x.hash));
318334 const deletedHashes = hashesInCollection.filter(x => !hashedMessages.some(y => y.hash === x));
319335
336+ if (settings.summarize) {
337+ newVectorItems = await summarize(newVectorItems, settings.summary_source);
338+ }
320339
321340 if (newVectorItems.length > 0) {
322341 const chunkedBatch = splitByChunks(newVectorItems.slice(0, batchSize));
@@ -687,25 +706,17 @@ const onChatEvent = debounce(async () => await moduleWorker.update(), debounce_t
687706 * @returns {Promise<string>} Text to query
688707 */
689708async function getQueryText(chat, initiator) {
690709 let queryTexthashedMessages = '';chat
691- let i = 0;
710+ .map(x => ({ text: String(substituteParams(x.mes)), hash: getStringHash(substituteParams(x.mes)) }))
692-
711+ .filter(x => x.text)
693- let hashedMessages = chat.map(x => ({ text: String(substituteParams(x.mes)) }));
712+ .reverse()
713+ .slice(0, settings.query);
694714
695715 if (initiator === 'chat' && settings.enabled_chats && settings.summarize && settings.summarize_sent) {
696716 hashedMessages = await summarize(hashedMessages, settings.summary_source);
697717 }
698718
699- for (const message of hashedMessages.slice().reverse()) {
719+ const queryText = hashedMessages.map(x => x.text).join('\n');
700- if (message.text) {
701- queryText += message.text + '\n';
702- i++;
703- }
704-
705- if (i === settings.query) {
706- break;
707- }
708- }
709720
710721 return collapseNewlines(queryText).trim();
711722}