Team Ai
Apppublic

KBaba7/llama.cpp

sourceHugging Faceapache-2.0updated 2y agoView on Hugging Face
0likes
tts.cpp974 linesDownload Raw Back to tts
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