Blame Raw
Cohee · 51ad27fb · · 481 lines (18.2 KB)
2 contributors
1import util from 'node:util';
2import { Buffer } from 'node:buffer';
3
4import fetch from 'node-fetch';
5import express from 'express';
6
7import { readSecret, SECRET_KEYS } from './secrets.js';
8import { readAllChunks, extractFileFromZipBuffer, forwardFetchResponse } from '../util.js';
9
10const API_NOVELAI = 'https://api.novelai.net';
11const TEXT_NOVELAI = 'https://text.novelai.net';
12const IMAGE_NOVELAI = 'https://image.novelai.net';
13
14// Constants for skip_cfg_above_sigma (Variety+) calculation
15const REFERENCE_PIXEL_COUNT = 1011712; // 832 * 1216 reference image size
16const SIGMA_MAGIC_NUMBER = 19; // Base sigma multiplier for V3 and V4 models
17const SIGMA_MAGIC_NUMBER_V4_5 = 58; // Base sigma multiplier for V4.5 models
18
19// Ban bracket generation, plus defaults
20const badWordsList = [
21 [3], [49356], [1431], [31715], [34387], [20765], [30702], [10691], [49333], [1266],
22 [19438], [43145], [26523], [41471], [2936], [85, 85], [49332], [7286], [1115], [24],
23];
24
25const eratoBadWordsList = [
26 [16067], [933, 11144], [25106, 11144], [58, 106901, 16073, 33710, 25, 109933],
27 [933, 58, 11144], [128030], [58, 30591, 33503, 17663, 100204, 25, 11144],
28];
29
30const hypeBotBadWordsList = [
31 [58], [60], [90], [92], [685], [1391], [1782], [2361], [3693], [4083], [4357], [4895],
32 [5512], [5974], [7131], [8183], [8351], [8762], [8964], [8973], [9063], [11208],
33 [11709], [11907], [11919], [12878], [12962], [13018], [13412], [14631], [14692],
34 [14980], [15090], [15437], [16151], [16410], [16589], [17241], [17414], [17635],
35 [17816], [17912], [18083], [18161], [18477], [19629], [19779], [19953], [20520],
36 [20598], [20662], [20740], [21476], [21737], [22133], [22241], [22345], [22935],
37 [23330], [23785], [23834], [23884], [25295], [25597], [25719], [25787], [25915],
38 [26076], [26358], [26398], [26894], [26933], [27007], [27422], [28013], [29164],
39 [29225], [29342], [29565], [29795], [30072], [30109], [30138], [30866], [31161],
40 [31478], [32092], [32239], [32509], [33116], [33250], [33761], [34171], [34758],
41 [34949], [35944], [36338], [36463], [36563], [36786], [36796], [36937], [37250],
42 [37913], [37981], [38165], [38362], [38381], [38430], [38892], [39850], [39893],
43 [41832], [41888], [42535], [42669], [42785], [42924], [43839], [44438], [44587],
44 [44926], [45144], [45297], [46110], [46570], [46581], [46956], [47175], [47182],
45 [47527], [47715], [48600], [48683], [48688], [48874], [48999], [49074], [49082],
46 [49146], [49946], [10221], [4841], [1427], [2602, 834], [29343], [37405], [35780], [2602], [50256],
47];
48
49// Used for phrase repetition penalty
50const repPenaltyAllowList = [
51 [49256, 49264, 49231, 49230, 49287, 85, 49255, 49399, 49262, 336, 333, 432, 363, 468, 492, 745, 401, 426, 623, 794,
52 1096, 2919, 2072, 7379, 1259, 2110, 620, 526, 487, 16562, 603, 805, 761, 2681, 942, 8917, 653, 3513, 506, 5301,
53 562, 5010, 614, 10942, 539, 2976, 462, 5189, 567, 2032, 123, 124, 125, 126, 127, 128, 129, 130, 131, 132, 588,
54 803, 1040, 49209, 4, 5, 6, 7, 8, 9, 10, 11, 12],
55];
56
57const eratoRepPenWhitelist = [
58 6, 1, 11, 13, 25, 198, 12, 9, 8, 279, 264, 459, 323, 477, 539, 912, 374, 574, 1051, 1550, 1587, 4536, 5828, 15058,
59 3287, 3250, 1461, 1077, 813, 11074, 872, 1202, 1436, 7846, 1288, 13434, 1053, 8434, 617, 9167, 1047, 19117, 706,
60 12775, 649, 4250, 527, 7784, 690, 2834, 15, 16, 17, 18, 19, 20, 21, 22, 23, 24, 1210, 1359, 608, 220, 596, 956,
61 3077, 44886, 4265, 3358, 2351, 2846, 311, 389, 315, 304, 520, 505, 430,
62];
63
64// Ban the dinkus and asterism
65const logitBiasExp = [
66 { 'sequence': [23], 'bias': -0.08, 'ensure_sequence_finish': false, 'generate_once': false },
67 { 'sequence': [21], 'bias': -0.08, 'ensure_sequence_finish': false, 'generate_once': false },
68];
69
70const eratoLogitBiasExp = [
71 { 'sequence': [12488], 'bias': -0.08, 'ensure_sequence_finish': false, 'generate_once': false },
72 { 'sequence': [128041], 'bias': -0.08, 'ensure_sequence_finish': false, 'generate_once': false },
73];
74
75function getBadWordsList(model) {
76 let list = [];
77
78 if (model.includes('hypebot')) {
79 list = hypeBotBadWordsList;
80 }
81
82 if (model.includes('clio') || model.includes('kayra')) {
83 list = badWordsList;
84 }
85
86 if (model.includes('erato')) {
87 list = eratoBadWordsList;
88 }
89
90 // Clone the list so we don't modify the original
91 return list.slice();
92}
93
94function getLogitBiasList(model) {
95 let list = [];
96
97 if (model.includes('erato')) {
98 list = eratoLogitBiasExp;
99 }
100
101 if (model.includes('clio') || model.includes('kayra')) {
102 list = logitBiasExp;
103 }
104
105 return list.slice();
106}
107
108function getRepPenaltyWhitelist(model) {
109 if (model.includes('clio') || model.includes('kayra')) {
110 return repPenaltyAllowList.flat();
111 }
112
113 if (model.includes('erato')) {
114 return eratoRepPenWhitelist.flat();
115 }
116
117 return null;
118}
119
120function calculateSkipCfgAboveSigma(width, height, modelName) {
121 const magicConstant = modelName?.includes('nai-diffusion-4-5')
122 ? SIGMA_MAGIC_NUMBER_V4_5
123 : SIGMA_MAGIC_NUMBER;
124
125 const pixelCount = width * height;
126 const ratio = pixelCount / REFERENCE_PIXEL_COUNT;
127
128 return Math.pow(ratio, 0.5) * magicConstant;
129}
130
131export const router = express.Router();
132
133router.post('/status', async function (req, res) {
134 if (!req.body) return res.sendStatus(400);
135 const api_key_novel = readSecret(req.user.directories, SECRET_KEYS.NOVEL);
136
137 if (!api_key_novel) {
138 console.warn('NovelAI Access Token is missing.');
139 return res.sendStatus(400);
140 }
141
142 try {
143 const response = await fetch(API_NOVELAI + '/user/subscription', {
144 method: 'GET',
145 headers: {
146 'Content-Type': 'application/json',
147 'Authorization': 'Bearer ' + api_key_novel,
148 },
149 });
150
151 if (response.ok) {
152 const data = await response.json();
153 return res.send(data);
154 } else if (response.status == 401) {
155 console.error('NovelAI Access Token is incorrect.');
156 return res.send({ error: true });
157 } else {
158 console.warn('NovelAI returned an error:', response.statusText);
159 return res.send({ error: true });
160 }
161 } catch (error) {
162 console.error(error);
163 return res.send({ error: true });
164 }
165});
166
167router.post('/generate', async function (req, res) {
168 if (!req.body) return res.sendStatus(400);
169
170 const api_key_novel = readSecret(req.user.directories, SECRET_KEYS.NOVEL);
171
172 if (!api_key_novel) {
173 console.warn('NovelAI Access Token is missing.');
174 return res.sendStatus(400);
175 }
176
177 const controller = new AbortController();
178 req.socket.removeAllListeners('close');
179 req.socket.on('close', function () {
180 controller.abort();
181 });
182
183 // Add customized bad words for Clio, Kayra, and Erato
184 const badWordsList = getBadWordsList(req.body.model);
185
186 if (Array.isArray(badWordsList) && Array.isArray(req.body.bad_words_ids)) {
187 for (const badWord of req.body.bad_words_ids) {
188 if (Array.isArray(badWord) && badWord.every(x => Number.isInteger(x))) {
189 badWordsList.push(badWord);
190 }
191 }
192 }
193
194 // Remove empty arrays from bad words list
195 for (const badWord of badWordsList) {
196 if (badWord.length === 0) {
197 badWordsList.splice(badWordsList.indexOf(badWord), 1);
198 }
199 }
200
201 // Add default biases for dinkus and asterism
202 const logitBiasList = getLogitBiasList(req.body.model);
203
204 if (Array.isArray(logitBiasList) && Array.isArray(req.body.logit_bias_exp)) {
205 logitBiasList.push(...req.body.logit_bias_exp);
206 }
207
208 const repPenWhitelist = getRepPenaltyWhitelist(req.body.model);
209
210 const data = {
211 'input': req.body.input,
212 'model': req.body.model,
213 'parameters': {
214 'use_string': req.body.use_string ?? true,
215 'temperature': req.body.temperature,
216 'max_length': req.body.max_length,
217 'min_length': req.body.min_length,
218 'tail_free_sampling': req.body.tail_free_sampling,
219 'repetition_penalty': req.body.repetition_penalty,
220 'repetition_penalty_range': req.body.repetition_penalty_range,
221 'repetition_penalty_slope': req.body.repetition_penalty_slope,
222 'repetition_penalty_frequency': req.body.repetition_penalty_frequency,
223 'repetition_penalty_presence': req.body.repetition_penalty_presence,
224 'repetition_penalty_whitelist': repPenWhitelist,
225 'top_a': req.body.top_a,
226 'top_p': req.body.top_p,
227 'top_k': req.body.top_k,
228 'typical_p': req.body.typical_p,
229 'mirostat_lr': req.body.mirostat_lr,
230 'mirostat_tau': req.body.mirostat_tau,
231 'phrase_rep_pen': req.body.phrase_rep_pen,
232 'stop_sequences': req.body.stop_sequences,
233 'bad_words_ids': badWordsList.length ? badWordsList : null,
234 'logit_bias_exp': logitBiasList,
235 'generate_until_sentence': req.body.generate_until_sentence,
236 'use_cache': req.body.use_cache,
237 'return_full_text': req.body.return_full_text,
238 'prefix': req.body.prefix,
239 'order': req.body.order,
240 'num_logprobs': req.body.num_logprobs,
241 'min_p': req.body.min_p,
242 'math1_temp': req.body.math1_temp,
243 'math1_quad': req.body.math1_quad,
244 'math1_quad_entropy_scale': req.body.math1_quad_entropy_scale,
245 },
246 };
247
248 // Tells the model to stop generation at '>'
249 if ('theme_textadventure' === req.body.prefix) {
250 if (req.body.model.includes('clio') || req.body.model.includes('kayra')) {
251 data.parameters.eos_token_id = 49405;
252 }
253 if (req.body.model.includes('erato')) {
254 data.parameters.eos_token_id = 29;
255 }
256 }
257
258 console.debug(util.inspect(data, { depth: 4 }));
259
260 const args = {
261 body: JSON.stringify(data),
262 headers: { 'Content-Type': 'application/json', 'Authorization': 'Bearer ' + api_key_novel },
263 signal: controller.signal,
264 };
265
266 try {
267 const baseURL = (req.body.model.includes('kayra') || req.body.model.includes('erato')) ? TEXT_NOVELAI : API_NOVELAI;
268 const url = req.body.streaming ? `${baseURL}/ai/generate-stream` : `${baseURL}/ai/generate`;
269 const response = await fetch(url, { method: 'POST', ...args });
270
271 if (req.body.streaming) {
272 // Pipe remote SSE stream to Express response
273 await forwardFetchResponse(response, res);
274 } else {
275 if (!response.ok) {
276 const text = await response.text();
277 let message = text;
278 console.warn(`Novel API returned error: ${response.status} ${response.statusText} ${text}`);
279
280 try {
281 const data = JSON.parse(text);
282 message = data.message;
283 } catch {
284 // ignore
285 }
286
287 return res.status(500).send({ error: { message } });
288 }
289
290 /** @type {any} */
291 const data = await response.json();
292 console.info('NovelAI Output', data?.output);
293 return res.send(data);
294 }
295 } catch (error) {
296 return res.send({ error: true });
297 }
298});
299
300router.post('/generate-image', async (request, response) => {
301 if (!request.body) {
302 return response.sendStatus(400);
303 }
304
305 const key = readSecret(request.user.directories, SECRET_KEYS.NOVEL);
306
307 if (!key) {
308 console.warn('NovelAI Access Token is missing.');
309 return response.sendStatus(400);
310 }
311
312 try {
313 console.debug('NAI Diffusion request:', request.body);
314 const generateUrl = `${IMAGE_NOVELAI}/ai/generate-image`;
315 const generateResult = await fetch(generateUrl, {
316 method: 'POST',
317 headers: {
318 'Authorization': `Bearer ${key}`,
319 'Content-Type': 'application/json',
320 },
321 body: JSON.stringify({
322 action: 'generate',
323 input: request.body.prompt ?? '',
324 model: request.body.model ?? 'nai-diffusion',
325 parameters: {
326 params_version: 3,
327 prefer_brownian: true,
328 negative_prompt: request.body.negative_prompt ?? '',
329 height: request.body.height ?? 512,
330 width: request.body.width ?? 512,
331 scale: request.body.scale ?? 9,
332 seed: request.body.seed >= 0 ? request.body.seed : Math.floor(Math.random() * 9999999999),
333 sampler: request.body.sampler ?? 'k_dpmpp_2m',
334 noise_schedule: request.body.scheduler ?? 'karras',
335 steps: request.body.steps ?? 28,
336 n_samples: 1,
337 // NAI handholding for prompts
338 ucPreset: 0,
339 qualityToggle: false,
340 add_original_image: false,
341 controlnet_strength: 1,
342 deliberate_euler_ancestral_bug: false,
343 dynamic_thresholding: request.body.decrisper ?? false,
344 legacy: false,
345 legacy_v3_extend: false,
346 sm: request.body.sm ?? false,
347 sm_dyn: request.body.sm_dyn ?? false,
348 uncond_scale: 1,
349 skip_cfg_above_sigma: request.body.variety_boost
350 ? calculateSkipCfgAboveSigma(
351 request.body.width ?? 512,
352 request.body.height ?? 512,
353 request.body.model ?? 'nai-diffusion',
354 )
355 : null,
356 use_coords: false,
357 characterPrompts: [],
358 reference_image_multiple: [],
359 reference_information_extracted_multiple: [],
360 reference_strength_multiple: [],
361 v4_negative_prompt: {
362 caption: {
363 base_caption: request.body.negative_prompt ?? '',
364 char_captions: [],
365 },
366 },
367 v4_prompt: {
368 caption: {
369 base_caption: request.body.prompt ?? '',
370 char_captions: [],
371 },
372 use_coords: false,
373 use_order: true,
374 },
375 },
376 }),
377 });
378
379 if (!generateResult.ok) {
380 const text = await generateResult.text();
381 console.warn('NovelAI returned an error.', generateResult.statusText, text);
382 return response.sendStatus(500);
383 }
384
385 const archiveBuffer = await generateResult.arrayBuffer();
386 const imageBuffer = await extractFileFromZipBuffer(archiveBuffer, '.png');
387
388 if (!imageBuffer) {
389 console.error('NovelAI generated an image, but the PNG file was not found.');
390 return response.sendStatus(500);
391 }
392
393 const originalBase64 = imageBuffer.toString('base64');
394
395 // No upscaling
396 if (isNaN(request.body.upscale_ratio) || request.body.upscale_ratio <= 1) {
397 return response.send(originalBase64);
398 }
399
400 try {
401 console.info('Upscaling image...');
402 const upscaleUrl = `${API_NOVELAI}/ai/upscale`;
403 const upscaleResult = await fetch(upscaleUrl, {
404 method: 'POST',
405 headers: {
406 'Authorization': `Bearer ${key}`,
407 'Content-Type': 'application/json',
408 },
409 body: JSON.stringify({
410 image: originalBase64,
411 height: request.body.height,
412 width: request.body.width,
413 scale: request.body.upscale_ratio,
414 }),
415 });
416
417 if (!upscaleResult.ok) {
418 const text = await upscaleResult.text();
419 throw new Error('NovelAI returned an error.', { cause: text });
420 }
421
422 const upscaledArchiveBuffer = await upscaleResult.arrayBuffer();
423 const upscaledImageBuffer = await extractFileFromZipBuffer(upscaledArchiveBuffer, '.png');
424
425 if (!upscaledImageBuffer) {
426 throw new Error('NovelAI upscaled an image, but the PNG file was not found.');
427 }
428
429 const upscaledBase64 = upscaledImageBuffer.toString('base64');
430
431 return response.send(upscaledBase64);
432 } catch (error) {
433 console.warn('NovelAI generated an image, but upscaling failed. Returning original image.', error);
434 return response.send(originalBase64);
435 }
436 } catch (error) {
437 console.error(error);
438 return response.sendStatus(500);
439 }
440});
441
442router.post('/generate-voice', async (request, response) => {
443 const token = readSecret(request.user.directories, SECRET_KEYS.NOVEL);
444
445 if (!token) {
446 console.error('NovelAI Access Token is missing.');
447 return response.sendStatus(400);
448 }
449
450 const text = request.body.text;
451 const voice = request.body.voice;
452
453 if (!text || !voice) {
454 return response.sendStatus(400);
455 }
456
457 try {
458 const url = `${API_NOVELAI}/ai/generate-voice?text=${encodeURIComponent(text)}&voice=-1&seed=${encodeURIComponent(voice)}&opus=false&version=v2`;
459 const result = await fetch(url, {
460 method: 'GET',
461 headers: {
462 'Authorization': `Bearer ${token}`,
463 'Accept': 'audio/mpeg',
464 },
465 });
466
467 if (!result.ok) {
468 const errorText = await result.text();
469 console.error('NovelAI returned an error.', result.statusText, errorText);
470 return response.sendStatus(500);
471 }
472
473 const chunks = await readAllChunks(result.body);
474 const buffer = Buffer.concat(chunks.map(chunk => new Uint8Array(chunk)));
475 response.setHeader('Content-Type', 'audio/mpeg');
476 return response.send(buffer);
477 } catch (error) {
478 console.error(error);
479 return response.sendStatus(500);
480 }
481});