Merge pull request #3213 from SillyTavern/expressions-webllm Expressions: Add WebLLM extension classification
Signed| @@ -15,6 +15,7 @@ import { SlashCommandEnumValue, enumTypes } from '../../slash-commands/SlashComm | |||
| 15 | import { commonEnumProviders } from '../../slash-commands/SlashCommandCommonEnumsProvider.js'; | 15 | import { commonEnumProviders } from '../../slash-commands/SlashCommandCommonEnumsProvider.js'; |
| 16 | import { slashCommandReturnHelper } from '../../slash-commands/SlashCommandReturnHelper.js'; | 16 | import { slashCommandReturnHelper } from '../../slash-commands/SlashCommandReturnHelper.js'; |
| 17 | import { SlashCommandClosure } from '../../slash-commands/SlashCommandClosure.js'; | 17 | import { SlashCommandClosure } from '../../slash-commands/SlashCommandClosure.js'; |
| 18 | import { generateWebLlmChatPrompt, isWebLlmSupported } from '../shared.js'; | ||
| 18 | export { MODULE_NAME }; | 19 | export { MODULE_NAME }; |
| 19 | 20 | ||
| 20 | const MODULE_NAME = 'expressions'; | 21 | const MODULE_NAME = 'expressions'; |
| @@ -59,6 +60,7 @@ const EXPRESSION_API = { | |||
| 59 | local: 0, | 60 | local: 0, |
| 60 | extras: 1, | 61 | extras: 1, |
| 61 | llm: 2, | 62 | llm: 2, |
| 63 | webllm: 3, | ||
| 62 | }; | 64 | }; |
| 63 | 65 | ||
| 64 | let expressionsList = null; | 66 | let expressionsList = null; |
| @@ -852,7 +854,7 @@ function setTalkingHeadState(newState) { | |||
| 852 | extension_settings.expressions.talkinghead = newState; // Store setting | 854 | extension_settings.expressions.talkinghead = newState; // Store setting |
| 853 | saveSettingsDebounced(); | 855 | saveSettingsDebounced(); |
| 854 | 856 | ||
| 855 | if (extension_settings.expressions.api == EXPRESSION_API.local || extension_settings.expressions.api == EXPRESSION_API.llm) { | 857 | if ([EXPRESSION_API.local, EXPRESSION_API.llm, EXPRESSION_API.webllm].includes(extension_settings.expressions.api)) { |
| 856 | return; | 858 | return; |
| 857 | } | 859 | } |
| 858 | 860 | ||
| @@ -1057,21 +1059,25 @@ function parseLlmResponse(emotionResponse, labels) { | |||
| 1057 | console.debug(`fuzzy search found: ${result[0].item} as closest for the LLM response:`, emotionResponse); | 1059 | console.debug(`fuzzy search found: ${result[0].item} as closest for the LLM response:`, emotionResponse); |
| 1058 | return result[0].item; | 1060 | return result[0].item; |
| 1059 | } | 1061 | } |
| 1062 | const lowerCaseResponse = String(emotionResponse || '').toLowerCase(); | ||
| 1063 | for (const label of labels) { | ||
| 1064 | if (lowerCaseResponse.includes(label.toLowerCase())) { | ||
| 1065 | console.debug(`Found label ${label} in the LLM response:`, emotionResponse); | ||
| 1066 | return label; | ||
| 1067 | } | ||
| 1068 | } | ||
| 1060 | } | 1069 | } |
| 1061 | 1070 | ||
| 1062 | throw new Error('Could not parse emotion response ' + emotionResponse); | 1071 | throw new Error('Could not parse emotion response ' + emotionResponse); |
| 1063 | } | 1072 | } |
| 1064 | 1073 | ||
| 1065 | function onTextGenSettingsReady(args) { | 1074 | /** |
| 1066 | // Only call if inside an API call | 1075 | * Gets the JSON schema for the LLM API. |
| 1067 | if (inApiCall && extension_settings.expressions.api === EXPRESSION_API.llm && isJsonSchemaSupported()) { | 1076 | * @param {string[]} emotions A list of emotions to search for. |
| 1068 | const emotions = DEFAULT_EXPRESSIONS.filter((e) => e != 'talkinghead'); | 1077 | * @returns {object} The JSON schema for the LLM API. |
| 1069 | Object.assign(args, { | 1078 | */ |
| 1070 | top_k: 1, | 1079 | function getJsonSchema(emotions) { |
| 1071 | stop: [], | 1080 | return { |
| 1072 | stopping_strings: [], | ||
| 1073 | custom_token_bans: [], | ||
| 1074 | json_schema: { | ||
| 1075 | $schema: 'http://json-schema.org/draft-04/schema#', | 1081 | $schema: 'http://json-schema.org/draft-04/schema#', |
| 1076 | type: 'object', | 1082 | type: 'object', |
| 1077 | properties: { | 1083 | properties: { |
| @@ -1083,7 +1089,19 @@ function onTextGenSettingsReady(args) { | |||
| 1083 | required: [ | 1089 | required: [ |
| 1084 | 'emotion', | 1090 | 'emotion', |
| 1085 | ], | 1091 | ], |
| 1086 | }, | 1092 | }; |
| 1093 | } | ||
| 1094 | |||
| 1095 | function onTextGenSettingsReady(args) { | ||
| 1096 | // Only call if inside an API call | ||
| 1097 | if (inApiCall && extension_settings.expressions.api === EXPRESSION_API.llm && isJsonSchemaSupported()) { | ||
| 1098 | const emotions = DEFAULT_EXPRESSIONS.filter((e) => e != 'talkinghead'); | ||
| 1099 | Object.assign(args, { | ||
| 1100 | top_k: 1, | ||
| 1101 | stop: [], | ||
| 1102 | stopping_strings: [], | ||
| 1103 | custom_token_bans: [], | ||
| 1104 | json_schema: getJsonSchema(emotions), | ||
| 1087 | }); | 1105 | }); |
| 1088 | } | 1106 | } |
| 1089 | } | 1107 | } |
| @@ -1139,6 +1157,22 @@ export async function getExpressionLabel(text, expressionsApi = extension_settin | |||
| 1139 | const emotionResponse = await generateRaw(text, main_api, false, false, prompt); | 1157 | const emotionResponse = await generateRaw(text, main_api, false, false, prompt); |
| 1140 | return parseLlmResponse(emotionResponse, expressionsList); | 1158 | return parseLlmResponse(emotionResponse, expressionsList); |
| 1141 | } | 1159 | } |
| 1160 | // Using WebLLM | ||
| 1161 | case EXPRESSION_API.webllm: { | ||
| 1162 | if (!isWebLlmSupported()) { | ||
| 1163 | console.warn('WebLLM is not supported. Using fallback expression'); | ||
| 1164 | return getFallbackExpression(); | ||
| 1165 | } | ||
| 1166 | |||
| 1167 | const expressionsList = await getExpressionsList(); | ||
| 1168 | const prompt = substituteParamsExtended(customPrompt, { labels: expressionsList }) || await getLlmPrompt(expressionsList); | ||
| 1169 | const messages = [ | ||
| 1170 | { role: 'user', content: text + '\n\n' + prompt }, | ||
| 1171 | ]; | ||
| 1172 | |||
| 1173 | const emotionResponse = await generateWebLlmChatPrompt(messages); | ||
| 1174 | return parseLlmResponse(emotionResponse, expressionsList); | ||
| 1175 | } | ||
| 1142 | // Extras | 1176 | // Extras |
| 1143 | default: { | 1177 | default: { |
| 1144 | const url = new URL(getApiUrl()); | 1178 | const url = new URL(getApiUrl()); |
| @@ -1603,7 +1637,7 @@ function onExpressionApiChanged() { | |||
| 1603 | const tempApi = this.value; | 1637 | const tempApi = this.value; |
| 1604 | if (tempApi) { | 1638 | if (tempApi) { |
| 1605 | extension_settings.expressions.api = Number(tempApi); | 1639 | extension_settings.expressions.api = Number(tempApi); |
| 1606 | $('.expression_llm_prompt_block').toggle(extension_settings.expressions.api === EXPRESSION_API.llm); | 1640 | $('.expression_llm_prompt_block').toggle([EXPRESSION_API.llm, EXPRESSION_API.webllm].includes(extension_settings.expressions.api)); |
| 1607 | expressionsList = null; | 1641 | expressionsList = null; |
| 1608 | spriteCache = {}; | 1642 | spriteCache = {}; |
| 1609 | moduleWorker(); | 1643 | moduleWorker(); |
| @@ -1940,7 +1974,7 @@ function migrateSettings() { | |||
| 1940 | 1974 | ||
| 1941 | await renderAdditionalExpressionSettings(); | 1975 | await renderAdditionalExpressionSettings(); |
| 1942 | $('#expression_api').val(extension_settings.expressions.api ?? EXPRESSION_API.extras); | 1976 | $('#expression_api').val(extension_settings.expressions.api ?? EXPRESSION_API.extras); |
| 1943 | $('.expression_llm_prompt_block').toggle(extension_settings.expressions.api === EXPRESSION_API.llm); | 1977 | $('.expression_llm_prompt_block').toggle([EXPRESSION_API.llm, EXPRESSION_API.webllm].includes(extension_settings.expressions.api)); |
| 1944 | $('#expression_llm_prompt').val(extension_settings.expressions.llmPrompt ?? ''); | 1978 | $('#expression_llm_prompt').val(extension_settings.expressions.llmPrompt ?? ''); |
| 1945 | $('#expression_llm_prompt').on('input', function () { | 1979 | $('#expression_llm_prompt').on('input', function () { |
| 1946 | extension_settings.expressions.llmPrompt = $(this).val(); | 1980 | extension_settings.expressions.llmPrompt = $(this).val(); |
| @@ -24,7 +24,8 @@ | |||
| 24 | <select id="expression_api" class="flex1 margin0"> | 24 | <select id="expression_api" class="flex1 margin0"> |
| 25 | <option value="0" data-i18n="Local">Local</option> | 25 | <option value="0" data-i18n="Local">Local</option> |
| 26 | <option value="1" data-i18n="Extras">Extras</option> | 26 | <option value="1" data-i18n="Extras">Extras</option> |
| 27 | <option value="2" data-i18n="LLM">LLM</option> | 27 | <option value="2" data-i18n="Main API">Main API</option> |
| 28 | <option value="3" data-i18n="WebLLM Extension">WebLLM Extension</option> | ||
| 28 | </select> | 29 | </select> |
| 29 | </div> | 30 | </div> |
| 30 | <div class="expression_llm_prompt_block m-b-1 m-t-1"> | 31 | <div class="expression_llm_prompt_block m-b-1 m-t-1"> |