KBaba7/llama.cpp
0
1#include "arg.h"2#include "common.h"3#include "sampling.h"4#include "log.h"5#include "llama.h"6 7#define _USE_MATH_DEFINES // For M_PI on MSVC8 9#include <algorithm>10#include <cmath>11#include <cstdio>12#include <fstream>13#include <map>14#include <regex>15#include <string>16#include <thread>17#include <vector>18 19//20// Terminal utils21//22 23#define SQR(X) ((X) * (X))24#define UNCUBE(x) x < 48 ? 0 : x < 115 ? 1 : (x - 35) / 4025 26/**27 * Quantizes 24-bit RGB to xterm256 code range [16,256).28 */29static int rgb2xterm256(int r, int g, int b) {30 unsigned char cube[] = {0, 0137, 0207, 0257, 0327, 0377};31 int av, ir, ig, ib, il, qr, qg, qb, ql;32 av = r * .299 + g * .587 + b * .114 + .5;33 ql = (il = av > 238 ? 23 : (av - 3) / 10) * 10 + 8;34 qr = cube[(ir = UNCUBE(r))];35 qg = cube[(ig = UNCUBE(g))];36 qb = cube[(ib = UNCUBE(b))];37 if (SQR(qr - r) + SQR(qg - g) + SQR(qb - b) <=38 SQR(ql - r) + SQR(ql - g) + SQR(ql - b))39 return ir * 36 + ig * 6 + ib + 020;40 return il + 0350;41}42 43static std::string set_xterm256_foreground(int r, int g, int b) {44 int x = rgb2xterm256(r, g, b);45 std::ostringstream oss;46 oss << "\033[38;5;" << x << "m";47 return oss.str();48}49 50const std::vector<std::string> k_colors = {51 set_xterm256_foreground(220, 5, 12),52 set_xterm256_foreground(232, 96, 28),53 set_xterm256_foreground(241, 147, 45),54 set_xterm256_foreground(246, 193, 65),55 set_xterm256_foreground(247, 240, 86),56 set_xterm256_foreground(144, 201, 135),57 set_xterm256_foreground( 78, 178, 101),58};59 60static void print_usage(int, char ** argv) {61 LOG("\nexample usage:\n");62 LOG("\n %s -m model.gguf -p \"Hello!\"\n", argv[0]);63 LOG("\n");64}65 66struct wav_header {67 char riff[4] = {'R', 'I', 'F', 'F'};68 uint32_t chunk_size;69 char wave[4] = {'W', 'A', 'V', 'E'};70 char fmt[4] = {'f', 'm', 't', ' '};71 uint32_t fmt_chunk_size = 16;72 uint16_t audio_format = 1; // PCM73 uint16_t num_channels = 1; // Mono74 uint32_t sample_rate;75 uint32_t byte_rate;76 uint16_t block_align;77 uint16_t bits_per_sample = 16;78 char data[4] = {'d', 'a', 't', 'a'};79 uint32_t data_size;80};81 82static void save_wav16(const std::string & fname, const std::vector<float> & data, int sample_rate) {83 std::ofstream file(fname, std::ios::binary);84 if (!file) {85 LOG_ERR("%s: Failed to open file '%s' for writing", __func__, fname.c_str());86 return;87 }88 89 wav_header header;90 header.sample_rate = sample_rate;91 header.byte_rate = header.sample_rate * header.num_channels * (header.bits_per_sample / 8);92 header.block_align = header.num_channels * (header.bits_per_sample / 8);93 header.data_size = data.size() * (header.bits_per_sample / 8);94 header.chunk_size = 36 + header.data_size;95 96 file.write(reinterpret_cast<const char*>(&header), sizeof(header));97 98 for (const auto & sample : data) {99 int16_t pcm_sample = static_cast<int16_t>(std::clamp(sample * 32767.0, -32768.0, 32767.0));100 file.write(reinterpret_cast<const char*>(&pcm_sample), sizeof(pcm_sample));101 }102 103 file.close();104}105 106static void fill_hann_window(int length, bool periodic, float * output) {107 int offset = -1;108 if (periodic) {109 offset = 0;110 }111 for (int i = 0; i < length; i++) {112 output[i] = 0.5 * (1.0 - cosf((2.0 * M_PI * i) / (length + offset)));113 }114}115 116// very poor-man fft117static void twiddle(float * real, float * imag, int k, int N) {118 float angle = 2 * M_PI * k / N;119 *real = cos(angle);120 *imag = sin(angle);121}122 123static void irfft(int n, const float * inp_cplx, float * out_real) {124 int N = n / 2 + 1;125 126 std::vector<float> real_input(N);127 std::vector<float> imag_input(N);128 for (int i = 0; i < N; ++i) {129 real_input[i] = inp_cplx[2 * i];130 imag_input[i] = inp_cplx[2 * i + 1];131 }132 133 std::vector<float> real_output(n);134 std::vector<float> imag_output(n);135 136 for (int k = 0; k < n; ++k) {137 real_output[k] = 0.0f;138 imag_output[k] = 0.0f;139 for (int m = 0; m < N; ++m) {140 float twiddle_real;141 float twiddle_imag;142 143 twiddle(&twiddle_real, &twiddle_imag, k * m, n);144 145 real_output[k] += real_input[m] * twiddle_real - imag_input[m] * twiddle_imag;146 imag_output[k] += real_input[m] * twiddle_imag + imag_input[m] * twiddle_real;147 }148 }149 150 for (int i = 0; i < n; ++i) {151 out_real[i] = real_output[i] / N;152 }153}154 155//156// y = torch.nn.functional.fold(157// data, output_size=(1, output_size), kernel_size=(1, self.win_length), stride=(1, self.hop_length),158// )[:, 0, 0, pad:-pad]159//160// data.shape = torch.Size([1, 1280, 261])161// output_size = 84480162// win_length = 1280163// hop_length = 320164// pad = 480165//166static void fold(const std::vector<float> & data, int64_t n_out, int64_t n_win, int64_t n_hop, int64_t n_pad, std::vector<float> & output) {167 int64_t output_height = n_out;168 int64_t kernel_w = n_win;169 int64_t stride_w = n_hop;170 int64_t width = n_out;171 172 output.resize(width, 0.0f);173 174 int64_t col_idx = 0;175 for (int64_t w_col = 0; w_col < width; ++w_col) {176 int64_t start = w_col * stride_w - n_pad;177 int64_t end = start + kernel_w;178 179 for (int64_t w_im = start; w_im < end; ++w_im) {180 if (w_im >= 0 && w_im < output_height && col_idx < (int64_t) data.size()) {181 output[w_im] += data[col_idx];182 }183 col_idx++;184 }185 }186 187 output.resize(n_out - 2 * n_pad);188}189 190// TODO: not optimized at all191static std::vector<float> embd_to_audio(192 const float * embd,193 const int n_codes,194 const int n_embd,195 const int n_thread) {196 const int n_fft = 1280;197 const int n_hop = 320;198 const int n_win = 1280;199 const int n_pad = (n_win - n_hop)/2;200 const int n_out = (n_codes - 1)*n_hop + n_win;201 202 std::vector<float> hann(n_fft);203 204 fill_hann_window(hann.size(), true, hann.data());205 206 int n_spec = n_embd*n_codes;207 208 std::vector<float> E (n_spec);209 std::vector<float> S (n_spec);210 std::vector<float> ST(n_spec);211 212 for (int l = 0; l < n_codes; ++l) {213 for (int k = 0; k < n_embd; ++k) {214 E[k*n_codes + l] = embd[l*n_embd + k];215 }216 }217 218 for (int k = 0; k < n_embd/2; ++k) {219 for (int l = 0; l < n_codes; ++l) {220 float mag = E[(k )*n_codes + l];221 float phi = E[(k + n_embd/2)*n_codes + l];222 223 mag = exp(mag);224 225 if (mag > 1e2) {226 mag = 1e2;227 }228 S[2*(k*n_codes + l) + 0] = mag*cosf(phi);229 S[2*(k*n_codes + l) + 1] = mag*sinf(phi);230 }231 }232 233 for (int l = 0; l < n_codes; ++l) {234 for (int k = 0; k < n_embd/2; ++k) {235 ST[l*n_embd + 2*k + 0] = S[2*(k*n_codes + l) + 0];236 ST[l*n_embd + 2*k + 1] = S[2*(k*n_codes + l) + 1];237 }238 }239 240 std::vector<float> res (n_codes*n_fft);241 std::vector<float> hann2(n_codes*n_fft);242 243 std::vector<std::thread> workers(n_thread);244 for (int i = 0; i < n_thread; ++i) {245 workers[i] = std::thread([&, i]() {246 for (int l = i; l < n_codes; l += n_thread) {247 irfft(n_fft, ST.data() + l*n_embd, res.data() + l*n_fft);248 for (int j = 0; j < n_fft; ++j) {249 res [l*n_fft + j] *= hann[j];250 hann2[l*n_fft + j] = hann[j] * hann[j];251 }252 }253 });254 }255 for (int i = 0; i < n_thread; ++i) {256 workers[i].join();257 }258 259 std::vector<float> audio;260 std::vector<float> env;261 262 fold(res, n_out, n_win, n_hop, n_pad, audio);263 fold(hann2, n_out, n_win, n_hop, n_pad, env); // TODO: can be done once264 265 for (size_t i = 0; i < audio.size(); ++i) {266 audio[i] /= env[i];267 }268 269 return audio;270}271 272static const std::map<int, std::string> ones = {273 {0, "zero"}, {1, "one"}, {2, "two"}, {3, "three"}, {4, "four"},274 {5, "five"}, {6, "six"}, {7, "seven"}, {8, "eight"}, {9, "nine"},275 {10, "ten"}, {11, "eleven"}, {12, "twelve"}, {13, "thirteen"}, {14, "fourteen"},276 {15, "fifteen"}, {16, "sixteen"}, {17, "seventeen"}, {18, "eighteen"}, {19, "nineteen"}277};278 279static const std::map<int, std::string> tens = {280 {2, "twenty"}, {3, "thirty"}, {4, "forty"}, {5, "fifty"},281 {6, "sixty"}, {7, "seventy"}, {8, "eighty"}, {9, "ninety"}282};283 284// Convert a number less than 1000 to words285static std::string convert_less_than_thousand(int num) {286 std::string result;287 288 if (num >= 100) {289 result += ones.at(num / 100) + " hundred ";290 num %= 100;291 }292 293 if (num >= 20) {294 result += tens.at(num / 10);295 if (num % 10 > 0) {296 result += "-" + ones.at(num % 10);297 }298 } else if (num > 0) {299 result += ones.at(num);300 }301 302 return result;303}304 305static std::string number_to_words(const std::string & number_str) {306 try {307 size_t decimal_pos = number_str.find('.');308 std::string integer_part = number_str.substr(0, decimal_pos);309 310 int int_number = std::stoi(integer_part);311 std::string result;312 313 if (int_number == 0) {314 result = "zero";315 } else {316 if (int_number >= 1000000000) {317 int billions = int_number / 1000000000;318 result += convert_less_than_thousand(billions) + " billion ";319 int_number %= 1000000000;320 }321 322 if (int_number >= 1000000) {323 int millions = int_number / 1000000;324 result += convert_less_than_thousand(millions) + " million ";325 int_number %= 1000000;326 }327 328 if (int_number >= 1000) {329 int thousands = int_number / 1000;330 result += convert_less_than_thousand(thousands) + " thousand ";331 int_number %= 1000;332 }333 334 if (int_number > 0) {335 result += convert_less_than_thousand(int_number);336 }337 }338 339 // Handle decimal part340 if (decimal_pos != std::string::npos) {341 result += " point";342 std::string decimal_part = number_str.substr(decimal_pos + 1);343 for (char digit : decimal_part) {344 result += " " + ones.at(digit - '0');345 }346 }347 348 return result;349 } catch (const std::exception& e) {350 // Skip if fails351 return " ";352 }353}354 355static std::string replace_numbers_with_words(const std::string & input_text) {356 std::regex number_pattern(R"(\d+(\.\d+)?)");357 std::string result;358 auto it = std::sregex_iterator(input_text.begin(), input_text.end(), number_pattern);359 auto end = std::sregex_iterator();360 361 size_t last_pos = 0;362 for (std::sregex_iterator i = it; i != end; ++i) {363 const std::smatch& match = *i;364 result.append(input_text, last_pos, match.position() - last_pos);365 result.append(number_to_words(match.str()));366 last_pos = match.position() + match.length();367 }368 result.append(input_text, last_pos);369 370 return result;371}372 373// Based on: https://github.com/edwko/OuteTTS/blob/a613e79c489d8256dd657ea9168d78de75895d82/outetts/version/v1/prompt_processor.py#L39374static std::string process_text(const std::string & text) {375 376 // For now I skipped text romanization as I am unsure how to handle377 // uroman and MeCab implementations in C++378 // maybe something like https://github.com/anyascii/anyascii/ could work.379 // currently only English would be supported in this function380 381 std::string processed_text = replace_numbers_with_words(text);382 383 std::transform(processed_text.begin(), processed_text.end(),384 processed_text.begin(), ::tolower);385 386 std::regex special_chars(R"([-_/,\.\\])");387 processed_text = std::regex_replace(processed_text, special_chars, " ");388 389 std::regex non_alpha(R"([^a-z\s])");390 processed_text = std::regex_replace(processed_text, non_alpha, "");391 392 std::regex multiple_spaces(R"(\s+)");393 processed_text = std::regex_replace(processed_text, multiple_spaces, " ");394 395 processed_text = std::regex_replace(processed_text, std::regex(R"(^\s+|\s+$)"), "");396 397 /*398 Replace spaces with the separator token same as in line 365399 400 for (auto & c : prompt_user) {401 if (c == ' ') {402 prompt_clean += "<|text_sep|>";403 */404 processed_text = std::regex_replace(processed_text, std::regex(R"(\s)"), "<|text_sep|>");405 406 return processed_text;407}408 409static void prompt_add(llama_tokens & prompt, llama_token token) {410 prompt.push_back(token);411}412 413static void prompt_add(llama_tokens & prompt, const llama_tokens & tokens) {414 prompt.insert(prompt.end(), tokens.begin(), tokens.end());415}416 417static void prompt_add(llama_tokens & prompt, const llama_vocab * vocab, const std::string & txt, bool add_special, bool parse_special) {418 auto tmp = common_tokenize(vocab, txt, add_special, parse_special);419 prompt_add(prompt, tmp);420}421 422static void prompt_init(llama_tokens & prompt, const llama_vocab * vocab) {423 prompt.clear();424 425 prompt_add(prompt, vocab, "<|im_start|>\n", true, true);426}427 428static std::vector<llama_token> prepare_guide_tokens(const llama_vocab * vocab, const std::string & str) {429 const std::string& delimiter = "<|text_sep|>";430 431 std::vector<llama_token> result;432 size_t start = 0;433 size_t end = str.find(delimiter);434 435 //first token is always a newline, as it was not previously added436 result.push_back(common_tokenize(vocab, "\n", false, true)[0]);437 438 while (end != std::string::npos) {439 std::string current_word = str.substr(start, end - start);440 auto tmp = common_tokenize(vocab, current_word, false, true);441 result.push_back(tmp[0]);442 start = end + delimiter.length();443 end = str.find(delimiter, start);444 }445 446 // Add the last part447 std::string current_word = str.substr(start);448 auto tmp = common_tokenize(vocab, current_word, false, true);449 if (tmp.size() > 0) {450 result.push_back(tmp[0]);451 }452 return result;453}454 455int main(int argc, char ** argv) {456 common_params params;457 458 params.prompt = "";459 460 params.n_predict = 4096;461 params.n_batch = 8192;462 params.n_ctx = 8192;463 464 params.sampling.top_k = 4;465 params.sampling.samplers = { COMMON_SAMPLER_TYPE_TOP_K, };466 467 if (!common_params_parse(argc, argv, params, LLAMA_EXAMPLE_TTS, print_usage)) {468 return 1;469 }470 471 const int n_parallel = params.n_parallel;472 const int n_predict = params.n_predict;473 474 common_init();475 476 // init LLM477 478 llama_backend_init();479 llama_numa_init(params.numa);480 481 llama_model * model_ttc = NULL; // text-to-codes482 llama_model * model_cts = NULL; // codes-to-speech483 484 llama_context * ctx_ttc = NULL;485 llama_context * ctx_cts = NULL;486 487 common_init_result llama_init_ttc = common_init_from_params(params);488 489 model_ttc = llama_init_ttc.model.get();490 ctx_ttc = llama_init_ttc.context.get();491 492 const llama_vocab * vocab = llama_model_get_vocab(model_ttc);493 494 // TODO: refactor in a common struct495 params.model = params.vocoder.model;496 params.model_url = params.vocoder.model_url;497 params.hf_repo = params.vocoder.hf_repo;498 params.hf_file = params.vocoder.hf_file;499 500 params.embedding = true;501 502 common_init_result llama_init_cts = common_init_from_params(params);503 504 model_cts = llama_init_cts.model.get();505 ctx_cts = llama_init_cts.context.get();506 507 std::vector<common_sampler *> smpl(n_parallel);508 for (int i = 0; i < n_parallel; ++i) {509 params.sampling.no_perf = (i != 0);510 params.sampling.seed = params.sampling.seed + 1;511 512 smpl[i] = common_sampler_init(model_ttc, params.sampling);513 }514 515 LOG_INF("sampler seed: %u\n", common_sampler_get_seed(smpl[0]));516 LOG_INF("sampler params: \n%s\n", params.sampling.print().c_str());517 LOG_INF("sampler chain: %s\n", common_sampler_print(smpl[0]).c_str());518 519 LOG_INF("%s: loading done\n", __func__);520 521 const auto t_main_start = ggml_time_us();522 523 std::vector<llama_token> codes;524 std::vector<llama_token> guide_tokens;525 526 // process prompt and generate voice codes527 {528 LOG_INF("%s: constructing prompt ..\n", __func__);529 530 std::vector<llama_token> prompt_inp;531 532 prompt_init(prompt_inp, vocab);533 534 prompt_add(prompt_inp, vocab, "<|text_start|>the<|text_sep|>overall<|text_sep|>package<|text_sep|>from<|text_sep|>just<|text_sep|>two<|text_sep|>people<|text_sep|>is<|text_sep|>pretty<|text_sep|>remarkable<|text_sep|>sure<|text_sep|>i<|text_sep|>have<|text_sep|>some<|text_sep|>critiques<|text_sep|>about<|text_sep|>some<|text_sep|>of<|text_sep|>the<|text_sep|>gameplay<|text_sep|>aspects<|text_sep|>but<|text_sep|>its<|text_sep|>still<|text_sep|>really<|text_sep|>enjoyable<|text_sep|>and<|text_sep|>it<|text_sep|>looks<|text_sep|>lovely<|text_sep|>", false, true);535 536 // convert the input text into the necessary format expected by OuteTTS537 {538 std::string prompt_clean = process_text(params.prompt);539 if (params.vocoder.use_guide_tokens) {540 guide_tokens = prepare_guide_tokens(vocab, prompt_clean);541 }542 543 LOG_INF("%s: prompt: '%s'\n", __func__, prompt_clean.c_str());544 545 prompt_add(prompt_inp, vocab, prompt_clean, false, true);546 }547 548 prompt_add(prompt_inp, vocab, "<|text_end|>\n", false, true);549 550 // disabled to save time on tokenizing each time551 // TODO: load voices from the json files552#if 0553 const std::string voice_data = R"(<|audio_start|>554the<|t_0.08|><|code_start|><|257|><|740|><|636|><|913|><|788|><|1703|><|code_end|>555overall<|t_0.36|><|code_start|><|127|><|201|><|191|><|774|><|700|><|532|><|1056|><|557|><|798|><|298|><|1741|><|747|><|1662|><|1617|><|1702|><|1527|><|368|><|1588|><|1049|><|1008|><|1625|><|747|><|1576|><|728|><|1019|><|1696|><|1765|><|code_end|>556package<|t_0.56|><|code_start|><|935|><|584|><|1319|><|627|><|1016|><|1491|><|1344|><|1117|><|1526|><|1040|><|239|><|1435|><|951|><|498|><|723|><|1180|><|535|><|789|><|1649|><|1637|><|78|><|465|><|1668|><|901|><|595|><|1675|><|117|><|1009|><|1667|><|320|><|840|><|79|><|507|><|1762|><|1508|><|1228|><|1768|><|802|><|1450|><|1457|><|232|><|639|><|code_end|>557from<|t_0.19|><|code_start|><|604|><|782|><|1682|><|872|><|1532|><|1600|><|1036|><|1761|><|647|><|1554|><|1371|><|653|><|1595|><|950|><|code_end|>558just<|t_0.25|><|code_start|><|1782|><|1670|><|317|><|786|><|1748|><|631|><|599|><|1155|><|1364|><|1524|><|36|><|1591|><|889|><|1535|><|541|><|440|><|1532|><|50|><|870|><|code_end|>559two<|t_0.24|><|code_start|><|1681|><|1510|><|673|><|799|><|805|><|1342|><|330|><|519|><|62|><|640|><|1138|><|565|><|1552|><|1497|><|1552|><|572|><|1715|><|1732|><|code_end|>560people<|t_0.39|><|code_start|><|593|><|274|><|136|><|740|><|691|><|633|><|1484|><|1061|><|1138|><|1485|><|344|><|428|><|397|><|1562|><|645|><|917|><|1035|><|1449|><|1669|><|487|><|442|><|1484|><|1329|><|1832|><|1704|><|600|><|761|><|653|><|269|><|code_end|>561is<|t_0.16|><|code_start|><|566|><|583|><|1755|><|646|><|1337|><|709|><|802|><|1008|><|485|><|1583|><|652|><|10|><|code_end|>562pretty<|t_0.32|><|code_start|><|1818|><|1747|><|692|><|733|><|1010|><|534|><|406|><|1697|><|1053|><|1521|><|1355|><|1274|><|816|><|1398|><|211|><|1218|><|817|><|1472|><|1703|><|686|><|13|><|822|><|445|><|1068|><|code_end|>563remarkable<|t_0.68|><|code_start|><|230|><|1048|><|1705|><|355|><|706|><|1149|><|1535|><|1787|><|1356|><|1396|><|835|><|1583|><|486|><|1249|><|286|><|937|><|1076|><|1150|><|614|><|42|><|1058|><|705|><|681|><|798|><|934|><|490|><|514|><|1399|><|572|><|1446|><|1703|><|1346|><|1040|><|1426|><|1304|><|664|><|171|><|1530|><|625|><|64|><|1708|><|1830|><|1030|><|443|><|1509|><|1063|><|1605|><|1785|><|721|><|1440|><|923|><|code_end|>564sure<|t_0.36|><|code_start|><|792|><|1780|><|923|><|1640|><|265|><|261|><|1525|><|567|><|1491|><|1250|><|1730|><|362|><|919|><|1766|><|543|><|1|><|333|><|113|><|970|><|252|><|1606|><|133|><|302|><|1810|><|1046|><|1190|><|1675|><|code_end|>565i<|t_0.08|><|code_start|><|123|><|439|><|1074|><|705|><|1799|><|637|><|code_end|>566have<|t_0.16|><|code_start|><|1509|><|599|><|518|><|1170|><|552|><|1029|><|1267|><|864|><|419|><|143|><|1061|><|0|><|code_end|>567some<|t_0.16|><|code_start|><|619|><|400|><|1270|><|62|><|1370|><|1832|><|917|><|1661|><|167|><|269|><|1366|><|1508|><|code_end|>568critiques<|t_0.60|><|code_start|><|559|><|584|><|1163|><|1129|><|1313|><|1728|><|721|><|1146|><|1093|><|577|><|928|><|27|><|630|><|1080|><|1346|><|1337|><|320|><|1382|><|1175|><|1682|><|1556|><|990|><|1683|><|860|><|1721|><|110|><|786|><|376|><|1085|><|756|><|1523|><|234|><|1334|><|1506|><|1578|><|659|><|612|><|1108|><|1466|><|1647|><|308|><|1470|><|746|><|556|><|1061|><|code_end|>569about<|t_0.29|><|code_start|><|26|><|1649|><|545|><|1367|><|1263|><|1728|><|450|><|859|><|1434|><|497|><|1220|><|1285|><|179|><|755|><|1154|><|779|><|179|><|1229|><|1213|><|922|><|1774|><|1408|><|code_end|>570some<|t_0.23|><|code_start|><|986|><|28|><|1649|><|778|><|858|><|1519|><|1|><|18|><|26|><|1042|><|1174|><|1309|><|1499|><|1712|><|1692|><|1516|><|1574|><|code_end|>571of<|t_0.07|><|code_start|><|197|><|716|><|1039|><|1662|><|64|><|code_end|>572the<|t_0.08|><|code_start|><|1811|><|1568|><|569|><|886|><|1025|><|1374|><|code_end|>573gameplay<|t_0.48|><|code_start|><|1269|><|1092|><|933|><|1362|><|1762|><|1700|><|1675|><|215|><|781|><|1086|><|461|><|838|><|1022|><|759|><|649|><|1416|><|1004|><|551|><|909|><|787|><|343|><|830|><|1391|><|1040|><|1622|><|1779|><|1360|><|1231|><|1187|><|1317|><|76|><|997|><|989|><|978|><|737|><|189|><|code_end|>574aspects<|t_0.56|><|code_start|><|1423|><|797|><|1316|><|1222|><|147|><|719|><|1347|><|386|><|1390|><|1558|><|154|><|440|><|634|><|592|><|1097|><|1718|><|712|><|763|><|1118|><|1721|><|1311|><|868|><|580|><|362|><|1435|><|868|><|247|><|221|><|886|><|1145|><|1274|><|1284|><|457|><|1043|><|1459|><|1818|><|62|><|599|><|1035|><|62|><|1649|><|778|><|code_end|>575but<|t_0.20|><|code_start|><|780|><|1825|><|1681|><|1007|><|861|><|710|><|702|><|939|><|1669|><|1491|><|613|><|1739|><|823|><|1469|><|648|><|code_end|>576its<|t_0.09|><|code_start|><|92|><|688|><|1623|><|962|><|1670|><|527|><|599|><|code_end|>577still<|t_0.27|><|code_start|><|636|><|10|><|1217|><|344|><|713|><|957|><|823|><|154|><|1649|><|1286|><|508|><|214|><|1760|><|1250|><|456|><|1352|><|1368|><|921|><|615|><|5|><|code_end|>578really<|t_0.36|><|code_start|><|55|><|420|><|1008|><|1659|><|27|><|644|><|1266|><|617|><|761|><|1712|><|109|><|1465|><|1587|><|503|><|1541|><|619|><|197|><|1019|><|817|><|269|><|377|><|362|><|1381|><|507|><|1488|><|4|><|1695|><|code_end|>579enjoyable<|t_0.49|><|code_start|><|678|><|501|><|864|><|319|><|288|><|1472|><|1341|><|686|><|562|><|1463|><|619|><|1563|><|471|><|911|><|730|><|1811|><|1006|><|520|><|861|><|1274|><|125|><|1431|><|638|><|621|><|153|><|876|><|1770|><|437|><|987|><|1653|><|1109|><|898|><|1285|><|80|><|593|><|1709|><|843|><|code_end|>580and<|t_0.15|><|code_start|><|1285|><|987|><|303|><|1037|><|730|><|1164|><|502|><|120|><|1737|><|1655|><|1318|><|code_end|>581it<|t_0.09|><|code_start|><|848|><|1366|><|395|><|1601|><|1513|><|593|><|1302|><|code_end|>582looks<|t_0.27|><|code_start|><|1281|><|1266|><|1755|><|572|><|248|><|1751|><|1257|><|695|><|1380|><|457|><|659|><|585|><|1315|><|1105|><|1776|><|736|><|24|><|736|><|654|><|1027|><|code_end|>583lovely<|t_0.56|><|code_start|><|634|><|596|><|1766|><|1556|><|1306|><|1285|><|1481|><|1721|><|1123|><|438|><|1246|><|1251|><|795|><|659|><|1381|><|1658|><|217|><|1772|><|562|><|952|><|107|><|1129|><|1112|><|467|><|550|><|1079|><|840|><|1615|><|1469|><|1380|><|168|><|917|><|836|><|1827|><|437|><|583|><|67|><|595|><|1087|><|1646|><|1493|><|1677|><|code_end|>)";584 585 auto tmp = common_tokenize(vocab, voice_data, false, true);586 printf("\n\n");587 for (int i = 0; i < tmp.size(); ++i) {588 printf("%d, ", tmp[i]);589 }590 printf("\n\n");591#else592 prompt_add(prompt_inp, llama_tokens {593 151667, 198, 1782, 155780, 151669, 151929, 152412, 152308, 152585,594 152460, 153375, 151670, 198, 74455, 155808, 151669, 151799,595 151873, 151863, 152446, 152372, 152204, 152728, 152229, 152470,596 151970, 153413, 152419, 153334, 153289, 153374, 153199, 152040,597 153260, 152721, 152680, 153297, 152419, 153248, 152400, 152691,598 153368, 153437, 151670, 198, 1722, 155828, 151669, 152607,599 152256, 152991, 152299, 152688, 153163, 153016, 152789, 153198,600 152712, 151911, 153107, 152623, 152170, 152395, 152852, 152207,601 152461, 153321, 153309, 151750, 152137, 153340, 152573, 152267,602 153347, 151789, 152681, 153339, 151992, 152512, 151751, 152179,603 153434, 153180, 152900, 153440, 152474, 153122, 153129, 151904,604 152311, 151670, 198, 1499, 155791, 151669, 152276, 152454,605 153354, 152544, 153204, 153272, 152708, 153433, 152319, 153226,606 153043, 152325, 153267, 152622, 151670, 198, 4250, 155797,607 151669, 153454, 153342, 151989, 152458, 153420, 152303, 152271,608 152827, 153036, 153196, 151708, 153263, 152561, 153207, 152213,609 152112, 153204, 151722, 152542, 151670, 198, 19789, 155796,610 151669, 153353, 153182, 152345, 152471, 152477, 153014, 152002,611 152191, 151734, 152312, 152810, 152237, 153224, 153169, 153224,612 152244, 153387, 153404, 151670, 198, 16069, 155811, 151669,613 152265, 151946, 151808, 152412, 152363, 152305, 153156, 152733,614 152810, 153157, 152016, 152100, 152069, 153234, 152317, 152589,615 152707, 153121, 153341, 152159, 152114, 153156, 153001, 153504,616 153376, 152272, 152433, 152325, 151941, 151670, 198, 285,617 155788, 151669, 152238, 152255, 153427, 152318, 153009, 152381,618 152474, 152680, 152157, 153255, 152324, 151682, 151670, 198,619 32955, 155804, 151669, 153490, 153419, 152364, 152405, 152682,620 152206, 152078, 153369, 152725, 153193, 153027, 152946, 152488,621 153070, 151883, 152890, 152489, 153144, 153375, 152358, 151685,622 152494, 152117, 152740, 151670, 198, 37448, 480, 155840, 151669,623 151902, 152720, 153377, 152027, 152378, 152821, 153207, 153459,624 153028, 153068, 152507, 153255, 152158, 152921, 151958, 152609,625 152748, 152822, 152286, 151714, 152730, 152377, 152353, 152470,626 152606, 152162, 152186, 153071, 152244, 153118, 153375, 153018,627 152712, 153098, 152976, 152336, 151843, 153202, 152297, 151736,628 153380, 153502, 152702, 152115, 153181, 152735, 153277, 153457,629 152393, 153112, 152595, 151670, 198, 19098, 155808, 151669,630 152464, 153452, 152595, 153312, 151937, 151933, 153197, 152239,631 153163, 152922, 153402, 152034, 152591, 153438, 152215, 151673,632 152005, 151785, 152642, 151924, 153278, 151805, 151974, 153482,633 152718, 152862, 153347, 151670, 198, 72, 155780, 151669, 151795,634 152111, 152746, 152377, 153471, 152309, 151670, 198, 19016,635 155788, 151669, 153181, 152271, 152190, 152842, 152224, 152701,636 152939, 152536, 152091, 151815, 152733, 151672, 151670, 198,637 14689, 155788, 151669, 152291, 152072, 152942, 151734, 153042,638 153504, 152589, 153333, 151839, 151941, 153038, 153180, 151670,639 198, 36996, 8303, 155832, 151669, 152231, 152256, 152835,640 152801, 152985, 153400, 152393, 152818, 152765, 152249, 152600,641 151699, 152302, 152752, 153018, 153009, 151992, 153054, 152847,642 153354, 153228, 152662, 153355, 152532, 153393, 151782, 152458,643 152048, 152757, 152428, 153195, 151906, 153006, 153178, 153250,644 152331, 152284, 152780, 153138, 153319, 151980, 153142, 152418,645 152228, 152733, 151670, 198, 9096, 155801, 151669, 151698,646 153321, 152217, 153039, 152935, 153400, 152122, 152531, 153106,647 152169, 152892, 152957, 151851, 152427, 152826, 152451, 151851,648 152901, 152885, 152594, 153446, 153080, 151670, 198, 14689,649 155795, 151669, 152658, 151700, 153321, 152450, 152530, 153191,650 151673, 151690, 151698, 152714, 152846, 152981, 153171, 153384,651 153364, 153188, 153246, 151670, 198, 1055, 155779, 151669,652 151869, 152388, 152711, 153334, 151736, 151670, 198, 1782,653 155780, 151669, 153483, 153240, 152241, 152558, 152697, 153046,654 151670, 198, 5804, 1363, 155820, 151669, 152941, 152764, 152605,655 153034, 153434, 153372, 153347, 151887, 152453, 152758, 152133,656 152510, 152694, 152431, 152321, 153088, 152676, 152223, 152581,657 152459, 152015, 152502, 153063, 152712, 153294, 153451, 153032,658 152903, 152859, 152989, 151748, 152669, 152661, 152650, 152409,659 151861, 151670, 198, 300, 7973, 155828, 151669, 153095, 152469,660 152988, 152894, 151819, 152391, 153019, 152058, 153062, 153230,661 151826, 152112, 152306, 152264, 152769, 153390, 152384, 152435,662 152790, 153393, 152983, 152540, 152252, 152034, 153107, 152540,663 151919, 151893, 152558, 152817, 152946, 152956, 152129, 152715,664 153131, 153490, 151734, 152271, 152707, 151734, 153321, 152450,665 151670, 198, 8088, 155792, 151669, 152452, 153497, 153353,666 152679, 152533, 152382, 152374, 152611, 153341, 153163, 152285,667 153411, 152495, 153141, 152320, 151670, 198, 1199, 155781,668 151669, 151764, 152360, 153295, 152634, 153342, 152199, 152271,669 151670, 198, 43366, 155799, 151669, 152308, 151682, 152889,670 152016, 152385, 152629, 152495, 151826, 153321, 152958, 152180,671 151886, 153432, 152922, 152128, 153024, 153040, 152593, 152287,672 151677, 151670, 198, 53660, 155808, 151669, 151727, 152092,673 152680, 153331, 151699, 152316, 152938, 152289, 152433, 153384,674 151781, 153137, 153259, 152175, 153213, 152291, 151869, 152691,675 152489, 151941, 152049, 152034, 153053, 152179, 153160, 151676,676 153367, 151670, 198, 268, 4123, 480, 155821, 151669, 152350,677 152173, 152536, 151991, 151960, 153144, 153013, 152358, 152234,678 153135, 152291, 153235, 152143, 152583, 152402, 153483, 152678,679 152192, 152533, 152946, 151797, 153103, 152310, 152293, 151825,680 152548, 153442, 152109, 152659, 153325, 152781, 152570, 152957,681 151752, 152265, 153381, 152515, 151670, 198, 437, 155787,682 151669, 152957, 152659, 151975, 152709, 152402, 152836, 152174,683 151792, 153409, 153327, 152990, 151670, 198, 275, 155781,684 151669, 152520, 153038, 152067, 153273, 153185, 152265, 152974,685 151670, 198, 94273, 155799, 151669, 152953, 152938, 153427,686 152244, 151920, 153423, 152929, 152367, 153052, 152129, 152331,687 152257, 152987, 152777, 153448, 152408, 151696, 152408, 152326,688 152699, 151670, 198, 385, 16239, 155828, 151669, 152306, 152268,689 153438, 153228, 152978, 152957, 153153, 153393, 152795, 152110,690 152918, 152923, 152467, 152331, 153053, 153330, 151889, 153444,691 152234, 152624, 151779, 152801, 152784, 152139, 152222, 152751,692 152512, 153287, 153141, 153052, 151840, 152589, 152508, 153499,693 152109, 152255, 151739, 152267, 152759, 153318, 153165, 153349,694 151670,});695#endif696 697 // print the prompt token-by-token698 699 LOG("\n");700 701 for (auto id : prompt_inp) {702 LOG("%s", common_token_to_piece(ctx_ttc, id).c_str());703 }704 705 LOG_INF("%s: prompt size: %d\n", __func__, (int) prompt_inp.size());706 707 LOG("\n");708 709 // create a llama_batch710 // we use this object to submit token data for decoding711 llama_batch batch = llama_batch_init(std::max(prompt_inp.size(), (size_t) n_parallel), 0, n_parallel);712 713 std::vector<llama_seq_id> seq_ids(n_parallel, 0);714 for (int32_t i = 0; i < n_parallel; ++i) {715 seq_ids[i] = i;716 }717 718 // evaluate the initial prompt719 for (size_t i = 0; i < prompt_inp.size(); ++i) {720 common_batch_add(batch, prompt_inp[i], i, seq_ids, false);721 }722 GGML_ASSERT(batch.n_tokens == (int) prompt_inp.size());723 724 // llama_decode will output logits only for the last token of the prompt725 batch.logits[batch.n_tokens - 1] = true;726 727 if (llama_decode(ctx_ttc, batch) != 0) {728 LOG_ERR("%s: llama_decode() failed\n", __func__);729 return 1;730 }731 732 if (n_parallel > 1) {733 LOG_INF("\n\n%s: generating %d sequences ...\n", __func__, n_parallel);734 }735 736 llama_synchronize(ctx_ttc);737 738 LOG_INF("%s: time for prompt: %.3f ms\n\n", __func__, (ggml_time_us() - t_main_start) / 1000.0f);739 740 const auto t_dec_start = ggml_time_us();741 742 // main loop743 744 // remember the batch index of the last token for each parallel sequence745 // we need this to determine which logits to sample from746 std::vector<int32_t> i_batch(n_parallel, batch.n_tokens - 1);747 748 int n_past = batch.n_tokens;749 int n_decode = 0;750 751 bool next_token_uses_guide_token = true;752 753 while (n_decode <= n_predict) {754 // prepare the next batch755 common_batch_clear(batch);756 757 // sample the next token for each parallel sequence / stream758 for (int32_t i = 0; i < n_parallel; ++i) {759 if (i_batch[i] < 0) {760 // the stream has already finished761 continue;762 }763 764 llama_token new_token_id = common_sampler_sample(smpl[i], ctx_ttc, i_batch[i]);765 766 //guide tokens help prevent hallucinations by forcing the TTS to use the correct word767 if (!guide_tokens.empty() && next_token_uses_guide_token && !llama_vocab_is_control(vocab, new_token_id) && !llama_vocab_is_eog(vocab, new_token_id)) {768 llama_token guide_token = guide_tokens[0];769 guide_tokens.erase(guide_tokens.begin());770 new_token_id = guide_token; //ensure correct word fragment is used771 }772 773 //this is the token id that always precedes a new word774 next_token_uses_guide_token = (new_token_id == 198);775 776 common_sampler_accept(smpl[i], new_token_id, true);777 778 codes.push_back(new_token_id);779 780 const auto * cands = common_sampler_get_candidates(smpl[i]);781 782 // is it an end of generation? -> mark the stream as finished783 if (llama_vocab_is_eog(vocab, new_token_id) || n_decode == n_predict) {784 std::string reason;785 if (llama_vocab_is_eog(vocab, new_token_id)) {786 reason = "eos";787 } else {788 reason = "n_predict";789 }790 791 i_batch[i] = -1;792 793 LOG("\n");794 if (n_parallel > 1) {795 LOG_CNT("\n");796 LOG_INF("%s: stream %d finished at n_past = %d, reason = '%s'\n", __func__, i, n_past, reason.c_str());797 }798 799 continue;800 }801 802 {803 const float p = cands->data[cands->selected].p;804 805 const int col = std::max(0, std::min((int) k_colors.size() - 1, (int) ((3*p)*float(k_colors.size()))));806 807 LOG_CNT("%s%d%s", k_colors[col].c_str(), i, "\033[0m");808 //LOG_CNT("%d", i);809 }810 811 i_batch[i] = batch.n_tokens;812 813 // push this new token for next evaluation814 common_batch_add(batch, new_token_id, n_past, { i }, true);815 }816 817 // all streams are finished818 if (batch.n_tokens == 0) {819 break;820 }821 822 n_decode += 1;823 n_past += 1;824 825 // evaluate the current batch with the transformer model826 if (llama_decode(ctx_ttc, batch)) {827 LOG_ERR("%s : failed to eval, return code %d\n", __func__, 1);828 return 1;829 }830 }831 832 llama_batch_free(batch);833 834 LOG("\n");835 LOG_INF("%s: time for decoder: %.3f ms\n", __func__, (ggml_time_us() - t_dec_start) / 1000.0f);836 }837 838 common_perf_print(ctx_ttc, smpl[0]);839 840 //std::vector<llama_token> codes = {198, 88225, 155856, 151669, 152205,841 // 153064, 152537, 153421, 153209, 152524, 151689, 152993, 152438, 152695,842 // 153091, 152945, 152829, 152534, 152934, 153020, 151997, 152263, 153010,843 // 153146, 152399, 153208, 152496, 151793, 152848, 152263, 152571, 153286,844 // 152227, 153300, 152934, 152263, 153208, 152263, 152965, 152430, 152296,845 // 153146, 152920, 152376, 152556, 153363, 151775, 152044, 152972, 152690,846 // 153379, 152368, 152233, 153422, 152490, 151996, 152022, 151694, 152061,847 // 153238, 152539, 153356, 152640, 153021, 153123, 151962, 153094, 151670,848 // 198, 20339, 13189, 155824, 151669, 152070, 152007, 152910, 151683,849 // 152000, 152373, 152760, 152046, 151735, 152334, 152394, 153073, 152908,850 // 151856, 151953, 153247, 153293, 151903, 153480, 153168, 152478, 153359,851 // 153429, 151905, 151678, 152567, 152411, 152165, 152556, 153075, 153424,852 // 151993, 152999, 153078, 152151, 152088, 153389, 152484, 151874, 151670,853 // 198, 285, 155784, 151669, 152226, 152126, 152638, 153215, 151729,854 // 152959, 153479, 153059, 151838, 151670, 198, 1782, 155783, 151669,855 // 153288, 153055, 153314, 152497, 152962, 152741, 152076, 153253, 151670,856 // 198, 471, 16488, 155825, 151669, 152060, 152916, 151893, 153469, 152501,857 // 152080, 152743, 151932, 153161, 152096, 152761, 152698, 153401, 153242,858 // 153336, 152441, 152838, 153467, 152706, 153496, 153310, 152422, 153360,859 // 153115, 152763, 151998, 152373, 153450, 152554, 151968, 153323, 152055,860 // 152468, 153111, 153358, 152813, 152010, 151770, 152823, 152960, 151670,861 // 198, 22627, 155823, 151669, 152814, 152366, 153484, 152931, 153441,862 // 152164, 152877, 152915, 153463, 151692, 152911, 152747, 152776, 151831,863 // 153449, 151882, 152975, 152031, 152513, 153150, 152448, 152667, 153133,864 // 153189, 152619, 153466, 152054, 152106, 153119, 152277, 152439, 153109,865 // 152997, 152141, 153154, 153256, 153311, 151922, 151670, 198, 1055,866 // 155781, 151669, 152633, 151850, 153060, 153270, 152560, 153348, 152729,867 // 151670, 198, 25312, 155803, 151669, 152521, 153403, 152561, 153337,868 // 153383, 152199, 153493, 153326, 151830, 152254, 152248, 152349, 152153,869 // 153007, 151823, 153037, 152575, 152457, 152406, 152592, 153116, 153365,870 // 153456, 151670, 198, 88225, 155817, 151669, 153271, 151925, 152218,871 // 152418, 152253, 153140, 151903, 153151, 152626, 152338, 152647, 153464,872 // 152785, 152768, 151711, 152037, 152033, 151804, 152216, 151701, 151855,873 // 152348, 152995, 152955, 152905, 152342, 152340, 153391, 153453, 152418,874 // 153415, 151990, 153083, 152884, 151670, 198, 151668, 198, 151645};875 876 {877 const std::string inp_txt = common_detokenize(ctx_ttc, codes, true);878 879 LOG("\n");880 LOG_INF("codes: '%s'\n", inp_txt.c_str());881 LOG_INF("%s: codes size: %d\n", __func__, (int) codes.size());882 }883 884 // remove all non-audio tokens (i.e. < 151672 || > 155772)885 codes.erase(std::remove_if(codes.begin(), codes.end(), [](llama_token t) { return t < 151672 || t > 155772; }), codes.end());886 887 {888 const std::string inp_txt = common_detokenize(ctx_ttc, codes, true);889 LOG_INF("codes audio: '%s'\n", inp_txt.c_str());890 LOG_INF("%s: codes audio size: %d\n", __func__, (int) codes.size());891 }892 893 for (auto & token : codes) {894 token -= 151672;895 }896 897 const auto t_voc_start = ggml_time_us();898 899 const int n_codes = codes.size();900 901 llama_batch batch = llama_batch_init(n_codes, 0, 1);902 903 for (size_t i = 0; i < codes.size(); ++i) {904 common_batch_add(batch, codes[i], i, { 0 }, true); // TODO: all logits?905 }906 GGML_ASSERT(batch.n_tokens == n_codes);907 908 if (llama_decode(ctx_cts, batch) != 0) {909 LOG_ERR("%s: llama_decode() failed\n", __func__);910 return 1;911 }912 913 llama_synchronize(ctx_cts);914 915 LOG_INF("%s: time for vocoder: %.3f ms\n", __func__, (ggml_time_us() - t_voc_start) / 1000.0f);916 917 const auto t_spec_start = ggml_time_us();918 919#if 1920 // spectral operations921 const int n_embd = llama_model_n_embd(model_cts);922 const float * embd = llama_get_embeddings(ctx_cts);923 924 auto audio = embd_to_audio(embd, n_codes, n_embd, params.cpuparams.n_threads);925 926#else927 // read the spectrogram from a file for debugging purposes928 std::vector<float> audio;929 {930 std::ifstream fin("out.bin", std::ios::binary);931 if (!fin) {932 LOG_ERR("%s: failed to open file '%s'\n", __func__, "out.bin");933 return 1;934 }935 936 std::vector<float> embd;937 938 int n_codes;939 int n_embd;940 941 fin.read(reinterpret_cast<char *>(&n_codes), sizeof(int));942 fin.read(reinterpret_cast<char *>(&n_embd), sizeof(int));943 944 embd.resize(n_codes * n_embd);945 fin.read(reinterpret_cast<char *>(embd.data()), n_codes * n_embd * sizeof(float));946 fin.close();947 948 LOG_INF("%s: n_codes: %d, n_embd: %d\n", __func__, n_codes, n_embd);949 950 audio = embd_to_audio(embd.data(), n_codes, n_embd, params.cpuparams.n_threads);951 }952#endif953 954 const std::string fname = "output.wav";955 956 const int n_sr = 24000; // sampling rate957 958 // zero out first 0.25 seconds959 for (int i = 0; i < 24000/4; ++i) {960 audio[i] = 0.0f;961 }962 963 LOG_INF("%s: time for spectral ops: %.3f ms\n", __func__, (ggml_time_us() - t_spec_start) / 1000.0f);964 LOG_INF("%s: total time: %.3f ms\n", __func__, (ggml_time_us() - t_main_start) / 1000.0f);965 966 save_wav16(fname, audio, n_sr);967 968 LOG_INF("%s: audio written to file '%s'\n", __func__, fname.c_str());969 970 llama_backend_free();971 972 return 0;973}974 