Blame Raw
Cohee · 51ad27fb · · 1137 lines (38.1 KB)
2 contributors
1import fs from 'node:fs';
2import path from 'node:path';
3import { Buffer } from 'node:buffer';
4import zlib from 'node:zlib';
5import { promisify } from 'node:util';
6
7import express from 'express';
8import fetch from 'node-fetch';
9import { sync as writeFileAtomicSync } from 'write-file-atomic';
10
11import { Tokenizer } from '@agnai/web-tokenizers';
12import { SentencePieceProcessor } from '@agnai/sentencepiece-js';
13import tiktoken from 'tiktoken';
14
15import { convertClaudePrompt } from '../prompt-converters.js';
16import { TEXTGEN_TYPES } from '../constants.js';
17import { setAdditionalHeaders } from '../additional-headers.js';
18import { getConfigValue, isValidUrl } from '../util.js';
19
20/**
21 * @typedef { (req: import('express').Request, res: import('express').Response) => Promise<any> } TokenizationHandler
22 */
23
24/**
25 * @type {{[key: string]: import('tiktoken').Tiktoken}} Tokenizers cache
26 */
27const tokenizersCache = {};
28
29/**
30 * @type {string[]}
31 */
32export const TEXT_COMPLETION_MODELS = [
33 'gpt-3.5-turbo-instruct',
34 'gpt-3.5-turbo-instruct-0914',
35 'text-davinci-003',
36 'text-davinci-002',
37 'text-davinci-001',
38 'text-curie-001',
39 'text-babbage-001',
40 'text-ada-001',
41 'code-davinci-002',
42 'code-davinci-001',
43 'code-cushman-002',
44 'code-cushman-001',
45 'text-davinci-edit-001',
46 'code-davinci-edit-001',
47 'text-embedding-ada-002',
48 'text-similarity-davinci-001',
49 'text-similarity-curie-001',
50 'text-similarity-babbage-001',
51 'text-similarity-ada-001',
52 'text-search-davinci-doc-001',
53 'text-search-curie-doc-001',
54 'text-search-babbage-doc-001',
55 'text-search-ada-doc-001',
56 'code-search-babbage-code-001',
57 'code-search-ada-code-001',
58];
59
60const BYTES_PER_TOKEN = 3.35;
61const IS_DOWNLOAD_ALLOWED = getConfigValue('enableDownloadableTokenizers', true, 'boolean');
62const gunzip = promisify(zlib.gunzip);
63
64/**
65 * Guesstimates the token count for a string.
66 * @param {string} str String to tokenize.
67 * @returns {number} Token count.
68 */
69function guesstimate(str) {
70 const byteLength = Buffer.byteLength(str, 'utf8');
71 return Math.ceil(byteLength / BYTES_PER_TOKEN);
72}
73
74/**
75 * Gets a path to the tokenizer model. Downloads the model if it's a URL.
76 * @param {string} model Model URL or path
77 * @param {string|undefined} fallbackModel Fallback model path
78 * @returns {Promise<string>} Path to the tokenizer model
79 */
80async function getPathToTokenizer(model, fallbackModel) {
81 if (!isValidUrl(model)) {
82 return model;
83 }
84
85 try {
86 const url = new URL(model);
87
88 if (!['https:', 'http:'].includes(url.protocol)) {
89 throw new Error('Invalid URL protocol');
90 }
91
92 const fileName = url.pathname.split('/').pop();
93
94 if (!fileName) {
95 throw new Error('Failed to extract the file name from the URL');
96 }
97
98 const CACHE_PATH = path.join(globalThis.DATA_ROOT, '_cache');
99 if (!fs.existsSync(CACHE_PATH)) {
100 fs.mkdirSync(CACHE_PATH, { recursive: true });
101 }
102
103 // If an uncompressed version exists, return it
104 const isCompressed = path.extname(fileName) === '.gz';
105 const uncompressedName = path.basename(fileName, '.gz');
106 const uncompressedPath = path.join(CACHE_PATH, uncompressedName);
107 if (isCompressed && fs.existsSync(uncompressedPath)) {
108 return uncompressedPath;
109 }
110
111 const cachedFile = path.join(CACHE_PATH, fileName);
112 if (fs.existsSync(cachedFile)) {
113 // If the file was downloaded manually
114 if (isCompressed) {
115 const compressedBuffer = await fs.promises.readFile(cachedFile);
116 const decompressedBuffer = await gunzip(compressedBuffer);
117 writeFileAtomicSync(uncompressedPath, decompressedBuffer);
118 await fs.promises.unlink(cachedFile);
119 return uncompressedPath;
120 }
121 return cachedFile;
122 }
123
124 if (!IS_DOWNLOAD_ALLOWED) {
125 throw new Error('Downloading tokenizers is disabled, the model is not cached');
126 }
127
128 console.info('Downloading tokenizer model:', model);
129 const response = await fetch(model);
130 if (!response.ok) {
131 throw new Error(`Failed to fetch the model: ${response.status} ${response.statusText}`);
132 }
133
134 const arrayBuffer = await response.arrayBuffer();
135 if (isCompressed) {
136 const decompressedBuffer = await gunzip(arrayBuffer);
137 writeFileAtomicSync(uncompressedPath, decompressedBuffer);
138 return uncompressedPath;
139 }
140
141 writeFileAtomicSync(cachedFile, Buffer.from(arrayBuffer));
142 return cachedFile;
143 } catch (error) {
144 const getLastSegment = str => str?.split('/')?.pop() || '';
145 if (fallbackModel) {
146 console.error(`Could not get a tokenizer from ${getLastSegment(model)}. Reason: ${error.message}. Using a fallback model: ${getLastSegment(fallbackModel)}.`);
147 return fallbackModel;
148 }
149
150 throw new Error(`Failed to instantiate a tokenizer and fallback is not provided. Reason: ${error.message}`);
151 }
152}
153
154/**
155 * Sentencepiece tokenizer for tokenizing text.
156 */
157class SentencePieceTokenizer {
158 /**
159 * @type {import('@agnai/sentencepiece-js').SentencePieceProcessor} Sentencepiece tokenizer instance
160 */
161 #instance;
162 /**
163 * @type {string} Path to the tokenizer model
164 */
165 #model;
166 /**
167 * @type {string|undefined} Path to the fallback model
168 */
169 #fallbackModel;
170
171 /**
172 * Creates a new Sentencepiece tokenizer.
173 * @param {string} model Path to the tokenizer model
174 * @param {string} [fallbackModel] Path to the fallback model
175 */
176 constructor(model, fallbackModel) {
177 this.#model = model;
178 this.#fallbackModel = fallbackModel;
179 }
180
181 /**
182 * Gets the Sentencepiece tokenizer instance.
183 * @returns {Promise<import('@agnai/sentencepiece-js').SentencePieceProcessor|null>} Sentencepiece tokenizer instance
184 */
185 async get() {
186 if (this.#instance) {
187 return this.#instance;
188 }
189
190 try {
191 const pathToModel = await getPathToTokenizer(this.#model, this.#fallbackModel);
192 this.#instance = new SentencePieceProcessor();
193 await this.#instance.load(pathToModel);
194 console.info('Instantiated the tokenizer for', path.parse(pathToModel).name);
195 return this.#instance;
196 } catch (error) {
197 console.error('Sentencepiece tokenizer failed to load: ' + this.#model, error);
198 return null;
199 }
200 }
201}
202
203/**
204 * Web tokenizer for tokenizing text.
205 */
206class WebTokenizer {
207 /**
208 * @type {Tokenizer} Web tokenizer instance
209 */
210 #instance;
211 /**
212 * @type {string} Path to the tokenizer model
213 */
214 #model;
215 /**
216 * @type {string|undefined} Path to the fallback model
217 */
218 #fallbackModel;
219
220 /**
221 * Creates a new Web tokenizer.
222 * @param {string} model Path to the tokenizer model
223 * @param {string} [fallbackModel] Path to the fallback model
224 */
225 constructor(model, fallbackModel) {
226 this.#model = model;
227 this.#fallbackModel = fallbackModel;
228 }
229
230 /**
231 * Gets the Web tokenizer instance.
232 * @returns {Promise<Tokenizer|null>} Web tokenizer instance
233 */
234 async get() {
235 if (this.#instance) {
236 return this.#instance;
237 }
238
239 try {
240 const pathToModel = await getPathToTokenizer(this.#model, this.#fallbackModel);
241 const fileBuffer = await fs.promises.readFile(pathToModel);
242 this.#instance = await Tokenizer.fromJSON(fileBuffer);
243 console.info('Instantiated the tokenizer for', path.parse(pathToModel).name);
244 return this.#instance;
245 } catch (error) {
246 console.error('Web tokenizer failed to load: ' + this.#model, error);
247 return null;
248 }
249 }
250}
251
252const spp_llama = new SentencePieceTokenizer('src/tokenizers/llama.model');
253const spp_nerd = new SentencePieceTokenizer('src/tokenizers/nerdstash.model');
254const spp_nerd_v2 = new SentencePieceTokenizer('src/tokenizers/nerdstash_v2.model');
255const spp_mistral = new SentencePieceTokenizer('src/tokenizers/mistral.model');
256const spp_yi = new SentencePieceTokenizer('src/tokenizers/yi.model');
257const spp_gemma = new SentencePieceTokenizer('src/tokenizers/gemma.model');
258const spp_jamba = new SentencePieceTokenizer('src/tokenizers/jamba.model');
259const claude_tokenizer = new WebTokenizer('src/tokenizers/claude.json');
260const llama3_tokenizer = new WebTokenizer('src/tokenizers/llama3.json');
261const commandRTokenizer = new WebTokenizer('https://github.com/SillyTavern/SillyTavern-Tokenizers/raw/main/command-r.json.gz', 'src/tokenizers/llama3.json');
262const commandATokenizer = new WebTokenizer('https://github.com/SillyTavern/SillyTavern-Tokenizers/raw/main/command-a.json.gz', 'src/tokenizers/llama3.json');
263const qwen2Tokenizer = new WebTokenizer('https://github.com/SillyTavern/SillyTavern-Tokenizers/raw/main/qwen2.json.gz', 'src/tokenizers/llama3.json');
264const nemoTokenizer = new WebTokenizer('https://github.com/SillyTavern/SillyTavern-Tokenizers/raw/main/nemo.json.gz', 'src/tokenizers/llama3.json');
265const deepseekTokenizer = new WebTokenizer('https://github.com/SillyTavern/SillyTavern-Tokenizers/raw/main/deepseek.json.gz', 'src/tokenizers/llama3.json');
266
267export const sentencepieceTokenizers = [
268 'llama',
269 'nerdstash',
270 'nerdstash_v2',
271 'mistral',
272 'yi',
273 'gemma',
274 'jamba',
275];
276
277export const webTokenizers = [
278 'claude',
279 'llama3',
280 'command-r',
281 'command-a',
282 'qwen2',
283 'nemo',
284 'deepseek',
285];
286
287/**
288 * Gets the Sentencepiece tokenizer by the model name.
289 * @param {string} model Sentencepiece model name
290 * @returns {SentencePieceTokenizer|null} Sentencepiece tokenizer
291 */
292export function getSentencepiceTokenizer(model) {
293 if (model.includes('llama')) {
294 return spp_llama;
295 }
296
297 if (model.includes('nerdstash')) {
298 return spp_nerd;
299 }
300
301 if (model.includes('mistral')) {
302 return spp_mistral;
303 }
304
305 if (model.includes('nerdstash_v2')) {
306 return spp_nerd_v2;
307 }
308
309 if (model.includes('yi')) {
310 return spp_yi;
311 }
312
313 if (model.includes('gemma')) {
314 return spp_gemma;
315 }
316
317 if (model.includes('jamba')) {
318 return spp_jamba;
319 }
320
321 return null;
322}
323
324/**
325 * Gets the Web tokenizer by the model name.
326 * @param {string} model Web tokenizer model name
327 * @returns {WebTokenizer|null} Web tokenizer
328 */
329export function getWebTokenizer(model) {
330 if (model.includes('llama3')) {
331 return llama3_tokenizer;
332 }
333
334 if (model.includes('claude')) {
335 return claude_tokenizer;
336 }
337
338 if (model.includes('command-r')) {
339 return commandRTokenizer;
340 }
341
342 if (model.includes('command-a')) {
343 return commandATokenizer;
344 }
345
346 if (model.includes('qwen2')) {
347 return qwen2Tokenizer;
348 }
349
350 if (model.includes('nemo')) {
351 return nemoTokenizer;
352 }
353
354 if (model.includes('deepseek')) {
355 return deepseekTokenizer;
356 }
357
358 return null;
359}
360
361/**
362 * Counts the token ids for the given text using the Sentencepiece tokenizer.
363 * @param {SentencePieceTokenizer} tokenizer Sentencepiece tokenizer
364 * @param {string} text Text to tokenize
365 * @returns { Promise<{ids: number[], count: number}> } Tokenization result
366 */
367async function countSentencepieceTokens(tokenizer, text) {
368 const instance = await tokenizer?.get();
369
370 // Fallback to strlen estimation
371 if (!instance) {
372 return {
373 ids: [],
374 count: guesstimate(text),
375 };
376 }
377
378 let cleaned = text; // cleanText(text); <-- cleaning text can result in an incorrect tokenization
379
380 let ids = instance.encodeIds(cleaned);
381 return {
382 ids,
383 count: ids.length,
384 };
385}
386
387/**
388 * Counts the tokens in the given array of objects using the Sentencepiece tokenizer.
389 * @param {SentencePieceTokenizer} tokenizer
390 * @param {object[]} array Array of objects to tokenize
391 * @returns {Promise<number>} Number of tokens
392 */
393async function countSentencepieceArrayTokens(tokenizer, array) {
394 const jsonBody = array.flatMap(x => Object.values(x)).join('\n\n');
395 const result = await countSentencepieceTokens(tokenizer, jsonBody);
396 const num_tokens = result.count;
397 return num_tokens;
398}
399
400async function getTiktokenChunks(tokenizer, ids) {
401 const decoder = new TextDecoder();
402 const chunks = [];
403
404 for (let i = 0; i < ids.length; i++) {
405 const id = ids[i];
406 const chunkTextBytes = await tokenizer.decode(new Uint32Array([id]));
407 const chunkText = decoder.decode(chunkTextBytes);
408 chunks.push(chunkText);
409 }
410
411 return chunks;
412}
413
414/**
415 * Gets the token chunks for the given token IDs using the Web tokenizer.
416 * @param {Tokenizer} tokenizer Web tokenizer instance
417 * @param {number[]} ids Token IDs
418 * @returns {string[]} Token chunks
419 */
420function getWebTokenizersChunks(tokenizer, ids) {
421 const chunks = [];
422
423 for (let i = 0, lastProcessed = 0; i < ids.length; i++) {
424 const chunkIds = ids.slice(lastProcessed, i + 1);
425 const chunkText = tokenizer.decode(new Int32Array(chunkIds));
426 if (chunkText === '�') {
427 continue;
428 }
429 chunks.push(chunkText);
430 lastProcessed = i + 1;
431 }
432
433 return chunks;
434}
435
436/**
437 * Gets the tokenizer model by the model name.
438 * @param {string} requestModel Models to use for tokenization
439 * @returns {string} Tokenizer model to use
440 */
441export function getTokenizerModel(requestModel) {
442 if (requestModel === 'o1' || requestModel.includes('o1-preview') || requestModel.includes('o1-mini') || requestModel.includes('o3-mini')) {
443 return 'o1';
444 }
445
446 if (requestModel.includes('gpt-5') || requestModel.includes('o3') || requestModel.includes('o4-mini')) {
447 return 'o1';
448 }
449
450 if (requestModel.includes('gpt-4o') || requestModel.includes('chatgpt-4o-latest')) {
451 return 'gpt-4o';
452 }
453
454 if (requestModel.includes('gpt-4.1') || requestModel.includes('gpt-4.5')) {
455 return 'gpt-4o';
456 }
457
458 if (requestModel.includes('gpt-4-32k')) {
459 return 'gpt-4-32k';
460 }
461
462 if (requestModel.includes('gpt-4')) {
463 return 'gpt-4';
464 }
465
466 if (requestModel.includes('gpt-3.5-turbo-0301')) {
467 return 'gpt-3.5-turbo-0301';
468 }
469
470 if (requestModel.includes('gpt-3.5-turbo')) {
471 return 'gpt-3.5-turbo';
472 }
473
474 if (TEXT_COMPLETION_MODELS.includes(requestModel)) {
475 return requestModel;
476 }
477
478 if (requestModel.includes('claude')) {
479 return 'claude';
480 }
481
482 if (requestModel.includes('llama3') || requestModel.includes('llama-3')) {
483 return 'llama3';
484 }
485
486 if (requestModel.includes('llama')) {
487 return 'llama';
488 }
489
490 if (requestModel.includes('mistral')) {
491 return 'mistral';
492 }
493
494 if (requestModel.includes('yi')) {
495 return 'yi';
496 }
497
498 if (requestModel.includes('deepseek')) {
499 return 'deepseek';
500 }
501
502 if (requestModel.includes('gemma') || requestModel.includes('gemini') || requestModel.includes('learnlm')) {
503 return 'gemma';
504 }
505
506 if (requestModel.includes('jamba')) {
507 return 'jamba';
508 }
509
510 if (requestModel.includes('qwen2')) {
511 return 'qwen2';
512 }
513
514 if (requestModel.includes('command-r')) {
515 return 'command-r';
516 }
517
518 if (requestModel.includes('command-a')) {
519 return 'command-a';
520 }
521
522 if (requestModel.includes('nemo')) {
523 return 'nemo';
524 }
525
526 // default
527 return 'gpt-3.5-turbo';
528}
529
530export function getTiktokenTokenizer(model) {
531 if (tokenizersCache[model]) {
532 return tokenizersCache[model];
533 }
534
535 const tokenizer = tiktoken.encoding_for_model(model);
536 console.info('Instantiated the tokenizer for', model);
537 tokenizersCache[model] = tokenizer;
538 return tokenizer;
539}
540
541/**
542 * Counts the tokens for the given messages using the WebTokenizer and Claude prompt conversion.
543 * @param {Tokenizer} tokenizer Web tokenizer
544 * @param {object[]} messages Array of messages
545 * @returns {number} Number of tokens
546 */
547export function countWebTokenizerTokens(tokenizer, messages) {
548 // Should be fine if we use the old conversion method instead of the messages API one i think?
549 const convertedPrompt = convertClaudePrompt(messages, false, '', false, false, '', false);
550
551 // Fallback to strlen estimation
552 if (!tokenizer) {
553 return guesstimate(convertedPrompt);
554 }
555
556 const count = tokenizer.encode(convertedPrompt).length;
557 return count;
558}
559
560/**
561 * Creates an API handler for encoding Sentencepiece tokens.
562 * @param {SentencePieceTokenizer} tokenizer Sentencepiece tokenizer
563 * @returns {TokenizationHandler} Handler function
564 */
565function createSentencepieceEncodingHandler(tokenizer) {
566 /**
567 * Request handler for encoding Sentencepiece tokens.
568 * @param {import('express').Request} request
569 * @param {import('express').Response} response
570 */
571 return async function (request, response) {
572 try {
573 if (!request.body) {
574 return response.sendStatus(400);
575 }
576
577 const text = request.body.text || '';
578 const instance = await tokenizer?.get();
579 const { ids, count } = await countSentencepieceTokens(tokenizer, text);
580 const chunks = instance?.encodePieces(text);
581 return response.send({ ids, count, chunks });
582 } catch (error) {
583 console.error(error);
584 return response.send({ ids: [], count: 0, chunks: [] });
585 }
586 };
587}
588
589/**
590 * Creates an API handler for decoding Sentencepiece tokens.
591 * @param {SentencePieceTokenizer} tokenizer Sentencepiece tokenizer
592 * @returns {TokenizationHandler} Handler function
593 */
594function createSentencepieceDecodingHandler(tokenizer) {
595 /**
596 * Request handler for decoding Sentencepiece tokens.
597 * @param {import('express').Request} request
598 * @param {import('express').Response} response
599 */
600 return async function (request, response) {
601 try {
602 if (!request.body) {
603 return response.sendStatus(400);
604 }
605
606 const ids = request.body.ids || [];
607 const instance = await tokenizer?.get();
608 if (!instance) throw new Error('Failed to load the Sentencepiece tokenizer');
609 const ops = ids.map(id => instance.decodeIds([id]));
610 const chunks = await Promise.all(ops);
611 const text = chunks.join('');
612 return response.send({ text, chunks });
613 } catch (error) {
614 console.error(error);
615 return response.send({ text: '', chunks: [] });
616 }
617 };
618}
619
620/**
621 * Creates an API handler for encoding Tiktoken tokens.
622 * @param {string} modelId Tiktoken model ID
623 * @returns {TokenizationHandler} Handler function
624 */
625function createTiktokenEncodingHandler(modelId) {
626 /**
627 * Request handler for encoding Tiktoken tokens.
628 * @param {import('express').Request} request
629 * @param {import('express').Response} response
630 */
631 return async function (request, response) {
632 try {
633 if (!request.body) {
634 return response.sendStatus(400);
635 }
636
637 const text = request.body.text || '';
638 const tokenizer = getTiktokenTokenizer(modelId);
639 const tokens = Object.values(tokenizer.encode(text));
640 const chunks = await getTiktokenChunks(tokenizer, tokens);
641 return response.send({ ids: tokens, count: tokens.length, chunks });
642 } catch (error) {
643 console.error(error);
644 return response.send({ ids: [], count: 0, chunks: [] });
645 }
646 };
647}
648
649/**
650 * Creates an API handler for decoding Tiktoken tokens.
651 * @param {string} modelId Tiktoken model ID
652 * @returns {TokenizationHandler} Handler function
653 */
654function createTiktokenDecodingHandler(modelId) {
655 /**
656 * Request handler for decoding Tiktoken tokens.
657 * @param {import('express').Request} request
658 * @param {import('express').Response} response
659 */
660 return async function (request, response) {
661 try {
662 if (!request.body) {
663 return response.sendStatus(400);
664 }
665
666 const ids = request.body.ids || [];
667 const tokenizer = getTiktokenTokenizer(modelId);
668 const textBytes = tokenizer.decode(new Uint32Array(ids));
669 const text = new TextDecoder().decode(textBytes);
670 return response.send({ text });
671 } catch (error) {
672 console.error(error);
673 return response.send({ text: '' });
674 }
675 };
676}
677
678/**
679 * Creates an API handler for encoding WebTokenizer tokens.
680 * @param {WebTokenizer} tokenizer WebTokenizer instance
681 * @returns {TokenizationHandler} Handler function
682 */
683function createWebTokenizerEncodingHandler(tokenizer) {
684 /**
685 * Request handler for encoding WebTokenizer tokens.
686 * @param {import('express').Request} request
687 * @param {import('express').Response} response
688 */
689 return async function (request, response) {
690 try {
691 if (!request.body) {
692 return response.sendStatus(400);
693 }
694
695 const text = request.body.text || '';
696 const instance = await tokenizer?.get();
697 if (!instance) throw new Error('Failed to load the Web tokenizer');
698 const tokens = Array.from(instance.encode(text));
699 const chunks = getWebTokenizersChunks(instance, tokens);
700 return response.send({ ids: tokens, count: tokens.length, chunks });
701 } catch (error) {
702 console.error(error);
703 return response.send({ ids: [], count: 0, chunks: [] });
704 }
705 };
706}
707
708/**
709 * Creates an API handler for decoding WebTokenizer tokens.
710 * @param {WebTokenizer} tokenizer WebTokenizer instance
711 * @returns {TokenizationHandler} Handler function
712 */
713function createWebTokenizerDecodingHandler(tokenizer) {
714 /**
715 * Request handler for decoding WebTokenizer tokens.
716 * @param {import('express').Request} request
717 * @param {import('express').Response} response
718 * @returns {Promise<any>}
719 */
720 return async function (request, response) {
721 try {
722 if (!request.body) {
723 return response.sendStatus(400);
724 }
725
726 const ids = request.body.ids || [];
727 const instance = await tokenizer?.get();
728 if (!instance) throw new Error('Failed to load the Web tokenizer');
729 const chunks = getWebTokenizersChunks(instance, ids);
730 const text = instance.decode(new Int32Array(ids));
731 return response.send({ text, chunks });
732 } catch (error) {
733 console.error(error);
734 return response.send({ text: '', chunks: [] });
735 }
736 };
737}
738
739export const router = express.Router();
740
741router.post('/llama/encode', createSentencepieceEncodingHandler(spp_llama));
742router.post('/nerdstash/encode', createSentencepieceEncodingHandler(spp_nerd));
743router.post('/nerdstash_v2/encode', createSentencepieceEncodingHandler(spp_nerd_v2));
744router.post('/mistral/encode', createSentencepieceEncodingHandler(spp_mistral));
745router.post('/yi/encode', createSentencepieceEncodingHandler(spp_yi));
746router.post('/gemma/encode', createSentencepieceEncodingHandler(spp_gemma));
747router.post('/jamba/encode', createSentencepieceEncodingHandler(spp_jamba));
748router.post('/gpt2/encode', createTiktokenEncodingHandler('gpt2'));
749router.post('/claude/encode', createWebTokenizerEncodingHandler(claude_tokenizer));
750router.post('/llama3/encode', createWebTokenizerEncodingHandler(llama3_tokenizer));
751router.post('/qwen2/encode', createWebTokenizerEncodingHandler(qwen2Tokenizer));
752router.post('/command-r/encode', createWebTokenizerEncodingHandler(commandRTokenizer));
753router.post('/command-a/encode', createWebTokenizerEncodingHandler(commandATokenizer));
754router.post('/nemo/encode', createWebTokenizerEncodingHandler(nemoTokenizer));
755router.post('/deepseek/encode', createWebTokenizerEncodingHandler(deepseekTokenizer));
756router.post('/llama/decode', createSentencepieceDecodingHandler(spp_llama));
757router.post('/nerdstash/decode', createSentencepieceDecodingHandler(spp_nerd));
758router.post('/nerdstash_v2/decode', createSentencepieceDecodingHandler(spp_nerd_v2));
759router.post('/mistral/decode', createSentencepieceDecodingHandler(spp_mistral));
760router.post('/yi/decode', createSentencepieceDecodingHandler(spp_yi));
761router.post('/gemma/decode', createSentencepieceDecodingHandler(spp_gemma));
762router.post('/jamba/decode', createSentencepieceDecodingHandler(spp_jamba));
763router.post('/gpt2/decode', createTiktokenDecodingHandler('gpt2'));
764router.post('/claude/decode', createWebTokenizerDecodingHandler(claude_tokenizer));
765router.post('/llama3/decode', createWebTokenizerDecodingHandler(llama3_tokenizer));
766router.post('/qwen2/decode', createWebTokenizerDecodingHandler(qwen2Tokenizer));
767router.post('/command-r/decode', createWebTokenizerDecodingHandler(commandRTokenizer));
768router.post('/command-a/decode', createWebTokenizerDecodingHandler(commandATokenizer));
769router.post('/nemo/decode', createWebTokenizerDecodingHandler(nemoTokenizer));
770router.post('/deepseek/decode', createWebTokenizerDecodingHandler(deepseekTokenizer));
771
772router.post('/openai/encode', async function (req, res) {
773 try {
774 const queryModel = String(req.query.model || '');
775
776 if (queryModel.includes('llama3') || queryModel.includes('llama-3')) {
777 const handler = createWebTokenizerEncodingHandler(llama3_tokenizer);
778 return handler(req, res);
779 }
780
781 if (queryModel.includes('llama')) {
782 const handler = createSentencepieceEncodingHandler(spp_llama);
783 return handler(req, res);
784 }
785
786 if (queryModel.includes('mistral')) {
787 const handler = createSentencepieceEncodingHandler(spp_mistral);
788 return handler(req, res);
789 }
790
791 if (queryModel.includes('yi')) {
792 const handler = createSentencepieceEncodingHandler(spp_yi);
793 return handler(req, res);
794 }
795
796 if (queryModel.includes('claude')) {
797 const handler = createWebTokenizerEncodingHandler(claude_tokenizer);
798 return handler(req, res);
799 }
800
801 if (queryModel.includes('gemma') || queryModel.includes('gemini')) {
802 const handler = createSentencepieceEncodingHandler(spp_gemma);
803 return handler(req, res);
804 }
805
806 if (queryModel.includes('jamba')) {
807 const handler = createSentencepieceEncodingHandler(spp_jamba);
808 return handler(req, res);
809 }
810
811 if (queryModel.includes('qwen2')) {
812 const handler = createWebTokenizerEncodingHandler(qwen2Tokenizer);
813 return handler(req, res);
814 }
815
816 if (queryModel.includes('command-r')) {
817 const handler = createWebTokenizerEncodingHandler(commandRTokenizer);
818 return handler(req, res);
819 }
820
821 if (queryModel.includes('command-a')) {
822 const handler = createWebTokenizerEncodingHandler(commandATokenizer);
823 return handler(req, res);
824 }
825
826 if (queryModel.includes('nemo')) {
827 const handler = createWebTokenizerEncodingHandler(nemoTokenizer);
828 return handler(req, res);
829 }
830
831 if (queryModel.includes('deepseek')) {
832 const handler = createWebTokenizerEncodingHandler(deepseekTokenizer);
833 return handler(req, res);
834 }
835
836 const model = getTokenizerModel(queryModel);
837 const handler = createTiktokenEncodingHandler(model);
838 return handler(req, res);
839 } catch (error) {
840 console.error(error);
841 return res.send({ ids: [], count: 0, chunks: [] });
842 }
843});
844
845router.post('/openai/decode', async function (req, res) {
846 try {
847 const queryModel = String(req.query.model || '');
848
849 if (queryModel.includes('llama3') || queryModel.includes('llama-3')) {
850 const handler = createWebTokenizerDecodingHandler(llama3_tokenizer);
851 return handler(req, res);
852 }
853
854 if (queryModel.includes('llama')) {
855 const handler = createSentencepieceDecodingHandler(spp_llama);
856 return handler(req, res);
857 }
858
859 if (queryModel.includes('mistral')) {
860 const handler = createSentencepieceDecodingHandler(spp_mistral);
861 return handler(req, res);
862 }
863
864 if (queryModel.includes('yi')) {
865 const handler = createSentencepieceDecodingHandler(spp_yi);
866 return handler(req, res);
867 }
868
869 if (queryModel.includes('claude')) {
870 const handler = createWebTokenizerDecodingHandler(claude_tokenizer);
871 return handler(req, res);
872 }
873
874 if (queryModel.includes('gemma') || queryModel.includes('gemini')) {
875 const handler = createSentencepieceDecodingHandler(spp_gemma);
876 return handler(req, res);
877 }
878
879 if (queryModel.includes('jamba')) {
880 const handler = createSentencepieceDecodingHandler(spp_jamba);
881 return handler(req, res);
882 }
883
884 if (queryModel.includes('qwen2')) {
885 const handler = createWebTokenizerDecodingHandler(qwen2Tokenizer);
886 return handler(req, res);
887 }
888
889 if (queryModel.includes('command-r')) {
890 const handler = createWebTokenizerDecodingHandler(commandRTokenizer);
891 return handler(req, res);
892 }
893
894 if (queryModel.includes('command-a')) {
895 const handler = createWebTokenizerDecodingHandler(commandATokenizer);
896 return handler(req, res);
897 }
898
899 if (queryModel.includes('nemo')) {
900 const handler = createWebTokenizerDecodingHandler(nemoTokenizer);
901 return handler(req, res);
902 }
903
904 if (queryModel.includes('deepseek')) {
905 const handler = createWebTokenizerDecodingHandler(deepseekTokenizer);
906 return handler(req, res);
907 }
908
909 const model = getTokenizerModel(queryModel);
910 const handler = createTiktokenDecodingHandler(model);
911 return handler(req, res);
912 } catch (error) {
913 console.error(error);
914 return res.send({ text: '' });
915 }
916});
917
918router.post('/openai/count', async function (req, res) {
919 try {
920 if (!req.body) return res.sendStatus(400);
921
922 let num_tokens = 0;
923 const queryModel = String(req.query.model || '');
924 const model = getTokenizerModel(queryModel);
925
926 if (model === 'claude') {
927 const instance = await claude_tokenizer.get();
928 if (!instance) throw new Error('Failed to load the Claude tokenizer');
929 num_tokens = countWebTokenizerTokens(instance, req.body);
930 return res.send({ 'token_count': num_tokens });
931 }
932
933 if (model === 'llama3' || model === 'llama-3') {
934 const instance = await llama3_tokenizer.get();
935 if (!instance) throw new Error('Failed to load the Llama3 tokenizer');
936 num_tokens = countWebTokenizerTokens(instance, req.body);
937 return res.send({ 'token_count': num_tokens });
938 }
939
940 if (model === 'llama') {
941 num_tokens = await countSentencepieceArrayTokens(spp_llama, req.body);
942 return res.send({ 'token_count': num_tokens });
943 }
944
945 if (model === 'mistral') {
946 num_tokens = await countSentencepieceArrayTokens(spp_mistral, req.body);
947 return res.send({ 'token_count': num_tokens });
948 }
949
950 if (model === 'yi') {
951 num_tokens = await countSentencepieceArrayTokens(spp_yi, req.body);
952 return res.send({ 'token_count': num_tokens });
953 }
954
955 if (model === 'gemma' || model === 'gemini') {
956 num_tokens = await countSentencepieceArrayTokens(spp_gemma, req.body);
957 return res.send({ 'token_count': num_tokens });
958 }
959
960 if (model === 'jamba') {
961 num_tokens = await countSentencepieceArrayTokens(spp_jamba, req.body);
962 return res.send({ 'token_count': num_tokens });
963 }
964
965 if (model === 'qwen2') {
966 const instance = await qwen2Tokenizer.get();
967 if (!instance) throw new Error('Failed to load the Qwen2 tokenizer');
968 num_tokens = countWebTokenizerTokens(instance, req.body);
969 return res.send({ 'token_count': num_tokens });
970 }
971
972 if (model === 'command-r') {
973 const instance = await commandRTokenizer.get();
974 if (!instance) throw new Error('Failed to load the Command-R tokenizer');
975 num_tokens = countWebTokenizerTokens(instance, req.body);
976 return res.send({ 'token_count': num_tokens });
977 }
978
979 if (model === 'command-a') {
980 const instance = await commandATokenizer.get();
981 if (!instance) throw new Error('Failed to load the Command-A tokenizer');
982 num_tokens = countWebTokenizerTokens(instance, req.body);
983 return res.send({ 'token_count': num_tokens });
984 }
985
986 if (model === 'nemo') {
987 const instance = await nemoTokenizer.get();
988 if (!instance) throw new Error('Failed to load the Nemo tokenizer');
989 num_tokens = countWebTokenizerTokens(instance, req.body);
990 return res.send({ 'token_count': num_tokens });
991 }
992
993 if (model === 'deepseek') {
994 const instance = await deepseekTokenizer.get();
995 if (!instance) throw new Error('Failed to load the DeepSeek tokenizer');
996 num_tokens = countWebTokenizerTokens(instance, req.body);
997 return res.send({ 'token_count': num_tokens });
998 }
999
1000 const tokensPerName = queryModel.includes('gpt-3.5-turbo-0301') ? -1 : 1;
1001 const tokensPerMessage = queryModel.includes('gpt-3.5-turbo-0301') ? 4 : 3;
1002 const tokensPadding = 3;
1003
1004 const tokenizer = getTiktokenTokenizer(model);
1005
1006 for (const msg of req.body) {
1007 try {
1008 num_tokens += tokensPerMessage;
1009 for (const [key, value] of Object.entries(msg)) {
1010 num_tokens += tokenizer.encode(value).length;
1011 if (key == 'name') {
1012 num_tokens += tokensPerName;
1013 }
1014 }
1015 } catch {
1016 console.warn('Error tokenizing message:', msg);
1017 }
1018 }
1019 num_tokens += tokensPadding;
1020
1021 // NB: Since 2023-10-14, the GPT-3.5 Turbo 0301 model shoves in 7-9 extra tokens to every message.
1022 // More details: https://community.openai.com/t/gpt-3-5-turbo-0301-showing-different-behavior-suddenly/431326/14
1023 if (queryModel.includes('gpt-3.5-turbo-0301')) {
1024 num_tokens += 9;
1025 }
1026
1027 // not needed for cached tokenizers
1028 //tokenizer.free();
1029
1030 res.send({ 'token_count': num_tokens });
1031 } catch (error) {
1032 console.error('An error counting tokens, using fallback estimation method', error);
1033 const jsonBody = JSON.stringify(req.body);
1034 const num_tokens = guesstimate(jsonBody);
1035 res.send({ 'token_count': num_tokens });
1036 }
1037});
1038
1039router.post('/remote/kobold/count', async function (request, response) {
1040 if (!request.body) {
1041 return response.sendStatus(400);
1042 }
1043 const text = String(request.body.text) || '';
1044 const baseUrl = String(request.body.url);
1045
1046 try {
1047 const args = {
1048 method: 'POST',
1049 body: JSON.stringify({ 'prompt': text }),
1050 headers: { 'Content-Type': 'application/json' },
1051 };
1052
1053 let url = String(baseUrl).replace(/\/$/, '');
1054 url += '/extra/tokencount';
1055
1056 const result = await fetch(url, args);
1057
1058 if (!result.ok) {
1059 console.warn(`API returned error: ${result.status} ${result.statusText}`);
1060 return response.send({ error: true });
1061 }
1062
1063 /** @type {any} */
1064 const data = await result.json();
1065 const count = data.value;
1066 const ids = data.ids ?? [];
1067 return response.send({ count, ids });
1068 } catch (error) {
1069 console.error(error);
1070 return response.send({ error: true });
1071 }
1072});
1073
1074router.post('/remote/textgenerationwebui/encode', async function (request, response) {
1075 if (!request.body) {
1076 return response.sendStatus(400);
1077 }
1078 const text = String(request.body.text) || '';
1079 const baseUrl = String(request.body.url);
1080 const model = String(request.body.model) || '';
1081
1082 try {
1083 const args = {
1084 method: 'POST',
1085 headers: { 'Content-Type': 'application/json' },
1086 };
1087
1088 setAdditionalHeaders(request, args, baseUrl);
1089
1090 // Convert to string + remove trailing slash + /v1 suffix
1091 let url = String(baseUrl).replace(/\/$/, '').replace(/\/v1$/, '');
1092
1093 switch (request.body.api_type) {
1094 case TEXTGEN_TYPES.TABBY:
1095 url += '/v1/token/encode';
1096 args.body = JSON.stringify({ 'text': text, 'add_bos_token': false, 'encode_special_tokens': false });
1097 break;
1098 case TEXTGEN_TYPES.KOBOLDCPP:
1099 url += '/api/extra/tokencount';
1100 args.body = JSON.stringify({ 'prompt': text, 'special': false });
1101 break;
1102 case TEXTGEN_TYPES.LLAMACPP:
1103 url += '/tokenize';
1104 args.body = JSON.stringify({ 'model': model, 'content': text });
1105 break;
1106 case TEXTGEN_TYPES.VLLM:
1107 url += '/tokenize';
1108 args.body = JSON.stringify({ 'model': model, 'prompt': text });
1109 break;
1110 case TEXTGEN_TYPES.APHRODITE:
1111 url += '/v1/tokenize';
1112 args.body = JSON.stringify({ 'model': model, 'prompt': text });
1113 break;
1114 default:
1115 url += '/v1/internal/encode';
1116 args.body = JSON.stringify({ 'text': text });
1117 break;
1118 }
1119
1120 const result = await fetch(url, args);
1121
1122 if (!result.ok) {
1123 console.warn(`API returned error: ${result.status} ${result.statusText}`);
1124 return response.send({ error: true });
1125 }
1126
1127 /** @type {any} */
1128 const data = await result.json();
1129 const count = (data?.length ?? data?.count ?? data?.value ?? data?.tokens?.length);
1130 const ids = (data?.tokens ?? data?.ids ?? []);
1131
1132 return response.send({ count, ids });
1133 } catch (error) {
1134 console.error(error);
1135 return response.send({ error: true });
1136 }
1137});