KBaba7/llama.cpp
0
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(¤t_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 