feature: derived templates This PR adds a simple hash based method for picking context and instruct templates based on the chat template, when provided by the back end.
Signed| @@ -267,6 +267,7 @@ import { applyBrowserFixes } from './scripts/browser-fixes.js'; | |||
| 267 | import { initServerHistory } from './scripts/server-history.js'; | 267 | import { initServerHistory } from './scripts/server-history.js'; |
| 268 | import { initSettingsSearch } from './scripts/setting-search.js'; | 268 | import { initSettingsSearch } from './scripts/setting-search.js'; |
| 269 | import { initBulkEdit } from './scripts/bulk-edit.js'; | 269 | import { initBulkEdit } from './scripts/bulk-edit.js'; |
| 270 | import { deriveTemplatesFromChatTemplate } from './scripts/chat-cemplates.js'; | ||
| 270 | 271 | ||
| 271 | //exporting functions and vars for mods | 272 | //exporting functions and vars for mods |
| 272 | export { | 273 | export { |
| @@ -1235,6 +1236,31 @@ async function getStatusTextgen() { | |||
| 1235 | const supportsTokenization = response.headers.get('x-supports-tokenization') === 'true'; | 1236 | const supportsTokenization = response.headers.get('x-supports-tokenization') === 'true'; |
| 1236 | supportsTokenization ? sessionStorage.setItem(TOKENIZER_SUPPORTED_KEY, 'true') : sessionStorage.removeItem(TOKENIZER_SUPPORTED_KEY); | 1237 | supportsTokenization ? sessionStorage.setItem(TOKENIZER_SUPPORTED_KEY, 'true') : sessionStorage.removeItem(TOKENIZER_SUPPORTED_KEY); |
| 1237 | 1238 | ||
| 1239 | const supportsChatTemplate = response.headers.get('x-supports-chat-template') === 'true'; | ||
| 1240 | |||
| 1241 | if (supportsChatTemplate) { | ||
| 1242 | const response = await fetch('/api/backends/text-completions/chat_template', { | ||
| 1243 | method: 'POST', | ||
| 1244 | headers: getRequestHeaders(), | ||
| 1245 | body: JSON.stringify({ | ||
| 1246 | api_server: endpoint, | ||
| 1247 | api_type: textgen_settings.type, | ||
| 1248 | }), | ||
| 1249 | }); | ||
| 1250 | |||
| 1251 | const data = await response.json(); | ||
| 1252 | if (data) { | ||
| 1253 | const chat_template = data.chat_template; | ||
| 1254 | console.log(`We have chat template ${chat_template.split('\n')[0]}...`); | ||
| 1255 | const templates = await deriveTemplatesFromChatTemplate(chat_template); | ||
| 1256 | if (templates) { | ||
| 1257 | const { context, instruct } = templates; | ||
| 1258 | selectContextPreset(context, { isAuto: true }); | ||
| 1259 | selectInstructPreset(instruct, { isAuto: true }); | ||
| 1260 | } | ||
| 1261 | } | ||
| 1262 | } | ||
| 1263 | |||
| 1238 | // We didn't get a 200 status code, but the endpoint has an explanation. Which means it DID connect, but I digress. | 1264 | // We didn't get a 200 status code, but the endpoint has an explanation. Which means it DID connect, but I digress. |
| 1239 | if (online_status === 'no_connection' && data.response) { | 1265 | if (online_status === 'no_connection' && data.response) { |
| 1240 | toastr.error(data.response, t`API Error`, { timeOut: 5000, preventDuplicates: true }); | 1266 | toastr.error(data.response, t`API Error`, { timeOut: 5000, preventDuplicates: true }); |
| @@ -0,0 +1,76 @@ | |||
| 1 | // https://developer.mozilla.org/en-US/docs/Web/API/SubtleCrypto/digest | ||
| 2 | async function digestMessage(message) { | ||
| 3 | const msgUint8 = new TextEncoder().encode(message); // encode as (utf-8) Uint8Array | ||
| 4 | const hashBuffer = await window.crypto.subtle.digest('SHA-256', msgUint8); // hash the message | ||
| 5 | const hashArray = Array.from(new Uint8Array(hashBuffer)); // convert buffer to byte array | ||
| 6 | const hashHex = hashArray | ||
| 7 | .map((b) => b.toString(16).padStart(2, '0')) | ||
| 8 | .join(''); // convert bytes to hex string | ||
| 9 | return hashHex; | ||
| 10 | } | ||
| 11 | |||
| 12 | // the hash can be obtained from command line e.g. via: MODEL=path_to_model; python -c "import json, hashlib, sys; print(hashlib.sha256(json.load(open('"$MODEL"/tokenizer_config.json'))['chat_template'].strip().encode()).hexdigest())" | ||
| 13 | // note that chat templates must be trimmed to match the llama.cpp metadata value | ||
| 14 | const derivations = { | ||
| 15 | // Meta | ||
| 16 | '93c0e9aa3629bbd77e68dbc0f5621f6e6b23aa8d74b932595cdb8d64684526d7': { | ||
| 17 | // Meta-Llama-3.1-8B-Instruct | ||
| 18 | // Meta-Llama-3.1-70B-Instruct | ||
| 19 | context: 'Llama 3 Instruct', | ||
| 20 | instruct: 'Llama 3 Instruct', | ||
| 21 | }, | ||
| 22 | 'd82792f95932f1c9cef5c4bd992f171225e3bf8c7b609b4557c9e1ec96be819f': { | ||
| 23 | // Llama-3.2-1B-Instruct | ||
| 24 | // Llama-3.2-3B-Instruct | ||
| 25 | context: 'Llama 3 Instruct', | ||
| 26 | instruct: 'Llama 3 Instruct', | ||
| 27 | }, | ||
| 28 | |||
| 29 | // Mistral | ||
| 30 | // Mistral Reference: https://github.com/mistralai/mistral-common | ||
| 31 | 'cafb64e0e9e5fd2503054b3479593fae39cbdfd52338ce8af9bb4664a8eb05bd': { | ||
| 32 | // Mistral-Small-Instruct-2409 | ||
| 33 | // Mistral-Large-Instruct-2407 | ||
| 34 | context: 'Mistral V2 & V3', | ||
| 35 | instruct: 'Mistral V2 & V3', | ||
| 36 | }, | ||
| 37 | '3c4ad5fa60dd8c7ccdf82fa4225864c903e107728fcaf859fa6052cb80c92ee9': { | ||
| 38 | // Mistral-Large-Instruct-2411 | ||
| 39 | context: 'Mistral V7', // https://huggingface.co/mistralai/Mistral-Large-Instruct-2411 | ||
| 40 | instruct: 'Mistral V7', | ||
| 41 | }, | ||
| 42 | 'e7deee034838db2bfc7487788a3013d8a307ab69f72f3c54a85f06fd76007d4e': { | ||
| 43 | // Mistral-Nemo-Instruct-2407 | ||
| 44 | context: 'Mistral V3-Tekken', | ||
| 45 | instruct: 'Mistral V3-Tekken', | ||
| 46 | }, | ||
| 47 | '26a59556925c987317ce5291811ba3b7f32ec4c647c400c6cc7e3a9993007ba7': { | ||
| 48 | // Mistral-7B-Instruct-v0.3 | ||
| 49 | context: 'Mistral V2 & V3', | ||
| 50 | instruct: 'Mistral V2 & V3', | ||
| 51 | }, | ||
| 52 | |||
| 53 | // Gemma | ||
| 54 | 'ecd6ae513fe103f0eb62e8ab5bfa8d0fe45c1074fa398b089c93a7e70c15cfd6': { | ||
| 55 | // gemma-2-9b-it | ||
| 56 | // gemma-2-27b-it | ||
| 57 | context: 'Gemma 2', | ||
| 58 | instruct: 'Gemma 2', | ||
| 59 | }, | ||
| 60 | |||
| 61 | // Cohere | ||
| 62 | '3b54f5c219ae1caa5c0bb2cdc7c001863ca6807cf888e4240e8739fa7eb9e02e': { | ||
| 63 | // command-r-08-2024 | ||
| 64 | context: 'Command R', | ||
| 65 | instruct: 'Command R', | ||
| 66 | }, | ||
| 67 | }; | ||
| 68 | |||
| 69 | export async function deriveTemplatesFromChatTemplate(chat_template) { | ||
| 70 | const hash = await digestMessage(chat_template); | ||
| 71 | if (hash in derivations) { | ||
| 72 | return derivations[hash]; | ||
| 73 | } | ||
| 74 | console.log(`Unknown chat template hash: ${hash}`); | ||
| 75 | return null; | ||
| 76 | } | ||
| @@ -218,6 +218,18 @@ router.post('/status', jsonParser, async function (request, response) { | |||
| 218 | } catch (error) { | 218 | } catch (error) { |
| 219 | console.error(`Failed to get TabbyAPI model info: ${error}`); | 219 | console.error(`Failed to get TabbyAPI model info: ${error}`); |
| 220 | } | 220 | } |
| 221 | } else if (apiType == TEXTGEN_TYPES.KOBOLDCPP) { | ||
| 222 | try { | ||
| 223 | const chatTemplateUrl = baseUrl + '/api/extra/chat_template'; | ||
| 224 | const chatTemplateReply = await fetch(chatTemplateUrl); | ||
| 225 | if (chatTemplateReply.ok) { | ||
| 226 | response.setHeader('x-supports-chat-template', 'true'); | ||
| 227 | } else { | ||
| 228 | console.log(`ct res = ${JSON.stringify(chatTemplateReply)}`); | ||
| 229 | } | ||
| 230 | } catch (error) { | ||
| 231 | console.error(`Failed to fetch chat template info: ${error}`); | ||
| 232 | } | ||
| 221 | } | 233 | } |
| 222 | 234 | ||
| 223 | return response.send({ result, data: data.data }); | 235 | return response.send({ result, data: data.data }); |
| @@ -227,6 +239,34 @@ router.post('/status', jsonParser, async function (request, response) { | |||
| 227 | } | 239 | } |
| 228 | }); | 240 | }); |
| 229 | 241 | ||
| 242 | router.post('/chat_template', jsonParser, async function (request, response) { | ||
| 243 | if (!request.body.api_server) return response.sendStatus(400); | ||
| 244 | |||
| 245 | try { | ||
| 246 | const baseUrl = trimV1(request.body.api_server); | ||
| 247 | const args = { | ||
| 248 | headers: { 'Content-Type': 'application/json' }, | ||
| 249 | }; | ||
| 250 | |||
| 251 | setAdditionalHeaders(request, args, baseUrl); | ||
| 252 | |||
| 253 | const chatTemplateUrl = baseUrl + '/api/extra/chat_template'; | ||
| 254 | const chatTemplateReply = await fetch(chatTemplateUrl, args); | ||
| 255 | |||
| 256 | if (!chatTemplateReply.ok) { | ||
| 257 | console.log('Chat template endpoint is offline.'); | ||
| 258 | return response.status(400); | ||
| 259 | } | ||
| 260 | |||
| 261 | /** @type {any} */ | ||
| 262 | const chatTemplate = await chatTemplateReply.json(); | ||
| 263 | return response.send(chatTemplate); | ||
| 264 | } catch (error) { | ||
| 265 | console.error(error); | ||
| 266 | return response.status(500); | ||
| 267 | } | ||
| 268 | }); | ||
| 269 | |||
| 230 | router.post('/generate', jsonParser, async function (request, response) { | 270 | router.post('/generate', jsonParser, async function (request, response) { |
| 231 | if (!request.body) return response.sendStatus(400); | 271 | if (!request.body) return response.sendStatus(400); |
| 232 | 272 | ||