This commit is contained in:
Georgi Gerganov
2026-09-18 09:36:15 +03:00
parent 1ec8188094
commit d093cf703e
+88 -24
View File
@@ -10,7 +10,10 @@
#include <algorithm>
#include <clocale>
#include <cmath>
#include <cstdint>
#include <cstdio>
#include <cstring>
#include <random>
#include <string>
#include <vector>
#include <ctime>
@@ -128,6 +131,9 @@ struct client {
std::string prompt;
std::string response;
std::vector<llama_token> prompt_tokens;
std::vector<llama_token> gen_tokens;
struct common_sampler * smpl = nullptr;
};
@@ -156,7 +162,9 @@ static std::vector<std::string> split_string(const std::string& input, char deli
int main(int argc, char ** argv) {
std::setlocale(LC_NUMERIC, "C");
srand(1234);
std::mt19937 rng(1234);
std::mt19937 token_rng(1234);
uint64_t logits_run_hash = 1469598103934665603ULL;
common_params params;
@@ -182,7 +190,7 @@ int main(int argc, char ** argv) {
const bool cont_batching = params.cont_batching;
// is the system prompt shared in the cache
const bool is_sp_shared = params.is_pp_shared;
bool is_sp_shared = params.is_pp_shared;
// extra text to insert in each client's prompt in order to make it larger
const int32_t n_junk = std::max(1, params.n_junk);
@@ -203,6 +211,10 @@ int main(int argc, char ** argv) {
auto * mem = llama_get_memory(ctx);
const llama_vocab * vocab = llama_model_get_vocab(model);
const bool no_vocab = llama_vocab_type(vocab) == LLAMA_VOCAB_TYPE_NONE;
if (no_vocab) {
is_sp_shared = false;
}
// load the prompts from an external file if there are any
if (params.prompt.empty()) {
@@ -244,7 +256,9 @@ int main(int argc, char ** argv) {
std::vector<llama_token> tokens_system;
tokens_system = common_tokenize(ctx, k_system, true);
if (!no_vocab) {
tokens_system = common_tokenize(ctx, k_system, true);
}
const int32_t n_tokens_system = tokens_system.size();
llama_seq_id g_seq_id = 0;
@@ -321,11 +335,9 @@ int main(int argc, char ** argv) {
client.t_start_prompt = ggml_time_us();
client.t_start_gen = 0;
client.input = k_prompts[rand() % k_prompts.size()];
client.input = k_prompts[rng() % k_prompts.size()];
client.response = "";
// construct the prompt:
// [system prompt] + [junk] + [user prompt]
client.n_past = 0;
client.prompt = "";
if (is_sp_shared) {
@@ -334,22 +346,45 @@ int main(int argc, char ** argv) {
client.prompt += k_system;
}
const int n_junk_cur = rand() % n_junk;
const int n_junk_cur = rng() % n_junk;
for (int i = 0; i < n_junk_cur; ++i) {
const int r = rand() % k_questions.size();
const int r = rng() % k_questions.size();
client.prompt += "User:\n" + k_questions[r] + "\nAssistant:\n " + k_answers[r] + "\n";
}
client.prompt += "User:\n" + client.input + "\nAssistant:\n";
common_sampler_reset(client.smpl);
// do not prepend BOS because we have a system prompt!
std::vector<llama_token> tokens_prompt;
tokens_prompt = common_tokenize(ctx, client.prompt, false);
if (no_vocab) {
const int32_t n_vocab = llama_vocab_n_tokens(vocab);
const int32_t n_prompt = 8;
const int32_t n_gen = params.n_predict > 0 ? params.n_predict : 128;
for (size_t i = 0; i < tokens_prompt.size(); ++i) {
common_batch_add(batch, tokens_prompt[i], client.n_past++, { client.id + 1 }, false);
client.prompt_tokens.resize(n_prompt);
for (auto & t : client.prompt_tokens) {
t = token_rng() % n_vocab;
}
client.gen_tokens.resize(n_gen);
for (auto & t : client.gen_tokens) {
t = token_rng() % n_vocab;
}
for (size_t i = 0; i < client.prompt_tokens.size(); ++i) {
common_batch_add(batch, client.prompt_tokens[i], client.n_past++, { client.id + 1 }, false);
}
client.n_prompt = n_prompt;
} else {
// do not prepend BOS because we have a system prompt!
std::vector<llama_token> tokens_prompt;
tokens_prompt = common_tokenize(ctx, client.prompt, false);
for (size_t i = 0; i < tokens_prompt.size(); ++i) {
common_batch_add(batch, tokens_prompt[i], client.n_past++, { client.id + 1 }, false);
}
client.n_prompt = tokens_prompt.size();
}
// extract the logits only for the last token
@@ -357,7 +392,6 @@ int main(int argc, char ** argv) {
batch.logits[batch.n_tokens - 1] = true;
}
client.n_prompt = tokens_prompt.size();
client.n_decoded = 0;
client.i_batch = batch.n_tokens - 1;
@@ -436,26 +470,56 @@ int main(int argc, char ** argv) {
//printf("client %d, seq %d, token %d, pos %d, batch %d\n",
// client.id, client.seq_id, client.sampled, client.n_decoded, client.i_batch);
const llama_token id = common_sampler_sample(client.smpl, ctx, client.i_batch - i);
const bool no_more_tokens = no_vocab &&
client.n_decoded >= (int32_t) client.gen_tokens.size();
common_sampler_accept(client.smpl, id, true);
const llama_token id = no_more_tokens
? 0
: (no_vocab
? client.gen_tokens[client.n_decoded]
: common_sampler_sample(client.smpl, ctx, client.i_batch - i));
if (client.n_decoded == 1) {
// start measuring generation time after the first token to make sure all concurrent clients
// have their prompt already processed
client.t_start_gen = ggml_time_us();
const float * logits = llama_get_logits_ith(ctx, client.i_batch - i);
const int32_t n_vocab = llama_vocab_n_tokens(vocab);
double logits_sum = 0.0;
uint64_t logits_hash = 1469598103934665603ULL;
for (int32_t j = 0; j < n_vocab; ++j) {
logits_sum += logits[j];
uint32_t u;
memcpy(&u, &logits[j], sizeof(u));
logits_hash ^= u;
logits_hash *= 1099511628211ULL;
}
const std::string token_str = common_token_to_piece(ctx, id);
logits_run_hash ^= logits_hash;
logits_run_hash *= 1099511628211ULL;
client.response += token_str;
client.sampled = id;
fprintf(stderr, "LOGITS_SUM client=%d seq=%d n_decoded=%d sum=%.9f hash=%llu run_hash=%llu\n",
client.id, client.seq_id, client.n_decoded, logits_sum,
(unsigned long long) logits_hash, (unsigned long long) logits_run_hash);
if (!no_more_tokens) {
if (!no_vocab) {
common_sampler_accept(client.smpl, id, true);
}
if (client.n_decoded == 1) {
// start measuring generation time after the first token to make sure all concurrent clients
// have their prompt already processed
client.t_start_gen = ggml_time_us();
}
const std::string token_str = no_vocab ? std::to_string(id) : common_token_to_piece(ctx, id);
client.response += token_str;
client.sampled = id;
}
//printf("client %d, seq %d, token %d, pos %d, batch %d: %s\n",
// client.id, client.seq_id, id, client.n_decoded, client.i_batch, token_str.c_str());
if (client.n_decoded > 2 &&
(llama_vocab_is_eog(vocab, id) ||
((!no_vocab && llama_vocab_is_eog(vocab, id)) ||
(params.n_predict > 0 && client.n_decoded >= params.n_predict) ||
client.response.find("User:") != std::string::npos)) {
// basic reverse prompt