Team Ai
Apppublic

KBaba7/llama.cpp

sourceHugging Faceapache-2.0updated 2y agoView on Hugging Face
0likes
parallel.cpp427 linesDownload Raw Back to parallel
1// A basic application simulating a server with multiple clients.2// The clients submit requests to the server and they are processed in parallel.3 4#include "arg.h"5#include "common.h"6#include "sampling.h"7#include "log.h"8#include "llama.h"9 10#include <cmath>11#include <cstdio>12#include <string>13#include <vector>14#include <ctime>15 16// trim whitespace from the beginning and end of a string17static std::string trim(const std::string & str) {18    size_t start = 0;19    size_t end = str.size();20 21    while (start < end && isspace(str[start])) {22        start += 1;23    }24 25    while (end > start && isspace(str[end - 1])) {26        end -= 1;27    }28 29    return str.substr(start, end - start);30}31 32static std::string k_system =33R"(Transcript of a never ending dialog, where the User interacts with an Assistant.34The Assistant is helpful, kind, honest, good at writing, and never fails to answer the User's requests immediately and with precision.35 36User: Recommend a nice restaurant in the area.37Assistant: I recommend the restaurant "The Golden Duck". It is a 5 star restaurant with a great view of the city. The food is delicious and the service is excellent. The prices are reasonable and the portions are generous. The restaurant is located at 123 Main Street, New York, NY 10001. The phone number is (212) 555-1234. The hours are Monday through Friday from 11:00 am to 10:00 pm. The restaurant is closed on Saturdays and Sundays.38User: Who is Richard Feynman?39Assistant: Richard Feynman was an American physicist who is best known for his work in quantum mechanics and particle physics. He was awarded the Nobel Prize in Physics in 1965 for his contributions to the development of quantum electrodynamics. He was a popular lecturer and author, and he wrote several books, including "Surely You're Joking, Mr. Feynman!" and "What Do You Care What Other People Think?".40User:)";41 42static std::vector<std::string> k_prompts = {43    "What is the meaning of life?",44    "Tell me an interesting fact about llamas.",45    "What is the best way to cook a steak?",46    "Are you familiar with the Special Theory of Relativity and can you explain it to me?",47    "Recommend some interesting books to read.",48    "What is the best way to learn a new language?",49    "How to get a job at Google?",50    "If you could have any superpower, what would it be?",51    "I want to learn how to play the piano.",52};53 54struct client {55    ~client() {56        if (smpl) {57            common_sampler_free(smpl);58        }59    }60 61    int32_t id = 0;62 63    llama_seq_id seq_id = -1;64 65    llama_token sampled;66 67    int64_t t_start_prompt;68    int64_t t_start_gen;69 70    int32_t n_prompt  = 0;71    int32_t n_decoded = 0;72    int32_t i_batch   = -1;73 74    std::string input;75    std::string prompt;76    std::string response;77 78    struct common_sampler * smpl = nullptr;79};80 81static void print_date_time() {82    std::time_t current_time = std::time(nullptr);83    std::tm* local_time = std::localtime(&current_time);84    char buffer[80];85    strftime(buffer, sizeof(buffer), "%Y-%m-%d %H:%M:%S", local_time);86 87    LOG_INF("\n");88    LOG_INF("\033[35mrun parameters as of %s\033[0m\n", buffer);89    LOG_INF("\n");90}91 92// Define a split string function to ...93static std::vector<std::string> split_string(const std::string& input, char delimiter) {94    std::vector<std::string> tokens;95    std::istringstream stream(input);96    std::string token;97    while (std::getline(stream, token, delimiter)) {98        tokens.push_back(token);99    }100    return tokens;101}102 103int main(int argc, char ** argv) {104    srand(1234);105 106    common_params params;107 108    if (!common_params_parse(argc, argv, params, LLAMA_EXAMPLE_PARALLEL)) {109        return 1;110    }111 112    common_init();113 114    // number of simultaneous "clients" to simulate115    const int32_t n_clients = params.n_parallel;116 117    // dedicate one sequence to the system prompt118    params.n_parallel += 1;119 120    // requests to simulate121    const int32_t n_seq = params.n_sequences;122 123    // insert new requests as soon as the previous one is done124    const bool cont_batching = params.cont_batching;125 126    const bool dump_kv_cache = params.dump_kv_cache;127 128    // init llama.cpp129    llama_backend_init();130    llama_numa_init(params.numa);131 132    // load the target model133    common_init_result llama_init = common_init_from_params(params);134 135    llama_model * model = llama_init.model.get();136    llama_context * ctx = llama_init.context.get();137 138    const llama_vocab * vocab = llama_model_get_vocab(model);139 140    // load the prompts from an external file if there are any141    if (params.prompt.empty()) {142        LOG_INF("\033[32mNo new questions so proceed with build-in defaults.\033[0m\n");143    } else {144        // Output each line of the input params.prompts vector and copy to k_prompts145        int index = 0;146        LOG_INF("\033[32mNow printing the external prompt file %s\033[0m\n\n", params.prompt_file.c_str());147 148        std::vector<std::string> prompts = split_string(params.prompt, '\n');149        for (const auto& prompt : prompts) {150            k_prompts.resize(index + 1);151            k_prompts[index] = prompt;152            index++;153            LOG_INF("%3d prompt: %s\n", index, prompt.c_str());154        }155    }156 157    LOG_INF("\n\n");158 159    const int n_ctx = llama_n_ctx(ctx);160 161    std::vector<client> clients(n_clients);162    for (size_t i = 0; i < clients.size(); ++i) {163        auto & client = clients[i];164        client.id = i;165        client.smpl = common_sampler_init(model, params.sampling);166    }167 168    std::vector<llama_token> tokens_system;169    tokens_system = common_tokenize(ctx, k_system, true);170    const int32_t n_tokens_system = tokens_system.size();171 172    llama_seq_id g_seq_id = 0;173 174    // the max batch size is as large as the context to handle cases where we get very long input prompt from multiple175    // users. regardless of the size, the main loop will chunk the batch into a maximum of params.n_batch tokens at a time176    llama_batch batch = llama_batch_init(n_ctx, 0, 1);177 178    int32_t n_total_prompt = 0;179    int32_t n_total_gen    = 0;180    int32_t n_cache_miss   = 0;181 182    struct llama_kv_cache_view kvc_view = llama_kv_cache_view_init(ctx, n_clients);183 184    const auto t_main_start = ggml_time_us();185 186    LOG_INF("%s: Simulating parallel requests from clients:\n", __func__);187    LOG_INF("%s: n_parallel = %d, n_sequences = %d, cont_batching = %d, system tokens = %d\n", __func__, n_clients, n_seq, cont_batching, n_tokens_system);188    LOG_INF("\n");189 190    {191        LOG_INF("%s: Evaluating the system prompt ...\n", __func__);192 193        for (int32_t i = 0; i < n_tokens_system; ++i) {194            common_batch_add(batch, tokens_system[i], i, { 0 }, false);195        }196 197        if (llama_decode(ctx, batch) != 0) {198            LOG_ERR("%s: llama_decode() failed\n", __func__);199            return 1;200        }201 202        // assign the system KV cache to all parallel sequences203        for (int32_t i = 1; i <= n_clients; ++i) {204            llama_kv_cache_seq_cp(ctx, 0, i, -1, -1);205        }206 207        LOG_INF("\n");208    }209 210    LOG_INF("Processing requests ...\n\n");211 212    while (true) {213        if (dump_kv_cache) {214            llama_kv_cache_view_update(ctx, &kvc_view);215            common_kv_cache_dump_view_seqs(kvc_view, 40);216        }217 218        common_batch_clear(batch);219 220        // decode any currently ongoing sequences221        for (auto & client : clients) {222            if (client.seq_id == -1) {223                continue;224            }225 226            client.i_batch = batch.n_tokens;227 228            common_batch_add(batch, client.sampled, n_tokens_system + client.n_prompt + client.n_decoded, { client.id + 1 }, true);229 230            client.n_decoded += 1;231        }232 233        if (batch.n_tokens == 0) {234            // all sequences have ended - clear the entire KV cache235            for (int i = 1; i <= n_clients; ++i) {236                llama_kv_cache_seq_rm(ctx, i, -1, -1);237                // but keep the system prompt238                llama_kv_cache_seq_cp(ctx, 0, i, -1, -1);239            }240 241            LOG_INF("%s: clearing the KV cache\n", __func__);242        }243 244        // insert new sequences for decoding245        if (cont_batching || batch.n_tokens == 0) {246            for (auto & client : clients) {247                if (client.seq_id == -1 && g_seq_id < n_seq) {248                    client.seq_id = g_seq_id;249 250                    client.t_start_prompt = ggml_time_us();251                    client.t_start_gen    = 0;252 253                    client.input    = k_prompts[rand() % k_prompts.size()];254                    client.prompt   = client.input + "\nAssistant:";255                    client.response = "";256 257                    common_sampler_reset(client.smpl);258 259                    // do not prepend BOS because we have a system prompt!260                    std::vector<llama_token> tokens_prompt;261                    tokens_prompt = common_tokenize(ctx, client.prompt, false);262 263                    for (size_t i = 0; i < tokens_prompt.size(); ++i) {264                        common_batch_add(batch, tokens_prompt[i], i + n_tokens_system, { client.id + 1 }, false);265                    }266 267                    // extract the logits only for the last token268                    if (batch.n_tokens > 0) {269                        batch.logits[batch.n_tokens - 1] = true;270                    }271 272                    client.n_prompt  = tokens_prompt.size();273                    client.n_decoded = 0;274                    client.i_batch   = batch.n_tokens - 1;275 276                    LOG_INF("\033[31mClient %3d, seq %4d, started decoding ...\033[0m\n", client.id, client.seq_id);277 278                    g_seq_id += 1;279 280                    // insert new requests one-by-one281                    //if (cont_batching) {282                    //    break;283                    //}284                }285            }286        }287 288        if (batch.n_tokens == 0) {289            break;290        }291 292        // process in chunks of params.n_batch293        int32_t n_batch = params.n_batch;294 295        for (int32_t i = 0; i < (int32_t) batch.n_tokens; i += n_batch) {296            // experiment: process in powers of 2297            //if (i + n_batch > (int32_t) batch.n_tokens && n_batch > 32) {298            //    n_batch /= 2;299            //    i -= n_batch;300            //    continue;301            //}302 303            const int32_t n_tokens = std::min(n_batch, (int32_t) (batch.n_tokens - i));304 305            llama_batch batch_view = {306                n_tokens,307                batch.token    + i,308                nullptr,309                batch.pos      + i,310                batch.n_seq_id + i,311                batch.seq_id   + i,312                batch.logits   + i,313            };314 315            const int ret = llama_decode(ctx, batch_view);316            if (ret != 0) {317                if (n_batch == 1 || ret < 0) {318                    // if you get here, it means the KV cache is full - try increasing it via the context size319                    LOG_ERR("%s : failed to decode the batch, n_batch = %d, ret = %d\n", __func__, n_batch, ret);320                    return 1;321                }322 323                LOG_ERR("%s : failed to decode the batch, retrying with n_batch = %d\n", __func__, n_batch / 2);324 325                n_cache_miss += 1;326 327                // retry with half the batch size to try to find a free slot in the KV cache328                n_batch /= 2;329                i -= n_batch;330 331                continue;332            }333 334            LOG_DBG("%s : decoded batch of %d tokens\n", __func__, n_tokens);335 336            for (auto & client : clients) {337                if (client.i_batch < (int) i || client.i_batch >= (int) (i + n_tokens)) {338                    continue;339                }340 341                //printf("client %d, seq %d, token %d, pos %d, batch %d\n",342                //        client.id, client.seq_id, client.sampled, client.n_decoded, client.i_batch);343 344                const llama_token id = common_sampler_sample(client.smpl, ctx, client.i_batch - i);345 346                common_sampler_accept(client.smpl, id, true);347 348                if (client.n_decoded == 1) {349                    // start measuring generation time after the first token to make sure all concurrent clients350                    // have their prompt already processed351                    client.t_start_gen = ggml_time_us();352                }353 354                const std::string token_str = common_token_to_piece(ctx, id);355 356                client.response += token_str;357                client.sampled = id;358 359                //printf("client %d, seq %d, token %d, pos %d, batch %d: %s\n",360                //        client.id, client.seq_id, id, client.n_decoded, client.i_batch, token_str.c_str());361 362                if (client.n_decoded > 2 &&363                        (llama_vocab_is_eog(vocab, id) ||364                         (params.n_predict > 0 && client.n_decoded + client.n_prompt >= params.n_predict) ||365                         client.response.find("User:") != std::string::npos ||366                         client.response.find('\n') != std::string::npos)) {367                    // basic reverse prompt368                    const size_t pos = client.response.find("User:");369                    if (pos != std::string::npos) {370                        client.response = client.response.substr(0, pos);371                    }372 373                    // delete only the generated part of the sequence, i.e. keep the system prompt in the cache374                    llama_kv_cache_seq_rm(ctx,    client.id + 1, -1, -1);375                    llama_kv_cache_seq_cp(ctx, 0, client.id + 1, -1, -1);376 377                    const auto t_main_end = ggml_time_us();378 379                    LOG_INF("\033[31mClient %3d, seq %3d/%3d, prompt %4d t, response %4d t, time %5.2f s, speed %5.2f t/s, cache miss %d \033[0m \n\nInput:    %s\n\033[35mResponse: %s\033[0m\n\n",380                            client.id, client.seq_id, n_seq, client.n_prompt, client.n_decoded,381                            (t_main_end - client.t_start_prompt) / 1e6,382                            (double) (client.n_prompt + client.n_decoded) / (t_main_end - client.t_start_prompt) * 1e6,383                            n_cache_miss,384                            ::trim(client.input).c_str(),385                            ::trim(client.response).c_str());386 387                    n_total_prompt += client.n_prompt;388                    n_total_gen    += client.n_decoded;389 390                    client.seq_id = -1;391                }392 393                client.i_batch = -1;394            }395        }396    }397 398    const auto t_main_end = ggml_time_us();399 400    print_date_time();401 402    LOG_INF("%s: n_parallel = %d, n_sequences = %d, cont_batching = %d, system tokens = %d\n", __func__, n_clients, n_seq, cont_batching, n_tokens_system);403    if (params.prompt_file.empty()) {404        params.prompt_file = "used built-in defaults";405    }406    LOG_INF("External prompt file: \033[32m%s\033[0m\n", params.prompt_file.c_str());407    LOG_INF("Model and path used:  \033[32m%s\033[0m\n\n", params.model.c_str());408 409    LOG_INF("Total prompt tokens: %6d, speed: %5.2f t/s\n", n_total_prompt, (double) (n_total_prompt              ) / (t_main_end - t_main_start) * 1e6);410    LOG_INF("Total gen tokens:    %6d, speed: %5.2f t/s\n", n_total_gen,    (double) (n_total_gen                 ) / (t_main_end - t_main_start) * 1e6);411    LOG_INF("Total speed (AVG):   %6s  speed: %5.2f t/s\n", "",             (double) (n_total_prompt + n_total_gen) / (t_main_end - t_main_start) * 1e6);412    LOG_INF("Cache misses:        %6d\n", n_cache_miss);413 414    LOG_INF("\n");415 416    // TODO: print sampling/grammar timings for all clients417    llama_perf_context_print(ctx);418 419    llama_batch_free(batch);420 421    llama_backend_free();422 423    LOG("\n\n");424 425    return 0;426}427