Files
research_and_development/stream/main.c
T
2025-07-12 12:22:05 -07:00

177 lines
4.7 KiB
C

#include "thirdparty/llama.cpp/ggml/include/ggml-backend.h"
#include <llama.h>
#include <stdio.h>
#include <stdlib.h>
#include <string.h>
typedef enum {
Unspecified,
Question,
Math,
Code,
Email,
Doc,
WebSearch,
} IntentType;
int main(int argc, char *argv[]) {
printf("Hello world\n");
char user_input[500];
printf("What's up, what can I help with?\n");
fgets(user_input, sizeof(user_input), stdin);
char prompt[1000];
snprintf(prompt, sizeof(prompt), "<|system|>You are a helpful assistant.<|end|>\n<|user|>%s<|end|>\n<|assistant|>", user_input);
// number of layers to offload to the GPU
int ngl = 99;
// number of tokens to predict
int n_predict = 1000;
// load dynamic backends
ggml_backend_load_all();
// initialize the model
struct llama_model_params model_params = llama_model_default_params();
model_params.n_gpu_layers = ngl;
struct llama_model *model = llama_model_load_from_file(
"Phi-3-mini-4k-instruct-q4.gguf", model_params);
if (model == NULL) {
fprintf(stderr, "%s: error: unable to load model\n", __func__);
return 1;
}
const struct llama_vocab *vocab = llama_model_get_vocab(model);
// tokenize the prompt
// find the number of tokens in the prompt
const int n_prompt =
-llama_tokenize(vocab, prompt, strlen(prompt), NULL, 0, true, true);
// allocate space for the tokens and tokenize the prompt
llama_token *prompt_tokens = malloc(n_prompt * sizeof(llama_token));
if (prompt_tokens == NULL) {
fprintf(stderr, "%s: error: failed to allocate memory for prompt tokens\n",
__func__);
return 1;
}
if (llama_tokenize(vocab, prompt, strlen(prompt), prompt_tokens, n_prompt,
true, true) < 0) {
fprintf(stderr, "%s: error: failed to tokenize the prompt\n", __func__);
return 1;
}
// initialize the context
struct llama_context_params ctx_params = llama_context_default_params();
// n_ctx is the context size
ctx_params.n_ctx = n_prompt + n_predict - 1;
// n_batch is the maximum number of tokens that can be processed in a single
// call to llama_decode
ctx_params.n_batch = n_prompt;
// enable performance counters
ctx_params.no_perf = false;
struct llama_context *ctx = llama_init_from_model(model, ctx_params);
if (ctx == NULL) {
fprintf(stderr, "%s: error: failed to create the llama_context\n",
__func__);
return 1;
}
// initialize the sampler
struct llama_sampler_chain_params sparams =
llama_sampler_chain_default_params();
sparams.no_perf = false;
struct llama_sampler *smpl = llama_sampler_chain_init(sparams);
llama_sampler_chain_add(smpl, llama_sampler_init_greedy());
// print the prompt token-by-token
for (int i = 0; i < n_prompt; i++) {
char buf[128];
int n = llama_token_to_piece(vocab, prompt_tokens[i], buf, sizeof(buf), 0,
true);
if (n < 0) {
fprintf(stderr, "%s: error: failed to convert token to piece\n",
__func__);
return 1;
}
buf[n] = '\0';
printf("%s", buf);
}
// prepare a batch for the prompt
llama_batch batch = llama_batch_get_one(prompt_tokens, n_prompt);
// main loop
const int64_t t_main_start = ggml_time_us();
int n_decode = 0;
llama_token new_token_id;
for (int n_pos = 0; n_pos + batch.n_tokens < n_prompt + n_predict;) {
// evaluate the current batch with the transformer model
if (llama_decode(ctx, batch)) {
fprintf(stderr, "%s : failed to eval, return code %d\n", __func__, 1);
return 1;
}
n_pos += batch.n_tokens;
// sample the next token
{
new_token_id = llama_sampler_sample(smpl, ctx, -1);
// is it an end of generation?
if (llama_vocab_is_eog(vocab, new_token_id)) {
break;
}
char buf[128];
int n =
llama_token_to_piece(vocab, new_token_id, buf, sizeof(buf), 0, true);
if (n < 0) {
fprintf(stderr, "%s: error: failed to convert token to piece\n",
__func__);
return 1;
}
buf[n] = '\0';
printf("%s", buf);
fflush(stdout);
// prepare the next batch with the sampled token
batch = llama_batch_get_one(&new_token_id, 1);
n_decode += 1;
}
}
printf("\n");
const int64_t t_main_end = ggml_time_us();
fprintf(stderr, "%s: decoded %d tokens in %.2f s, speed: %.2f t/s\n",
__func__, n_decode, (t_main_end - t_main_start) / 1000000.0f,
n_decode / ((t_main_end - t_main_start) / 1000000.0f));
fprintf(stderr, "\n");
llama_perf_sampler_print(smpl);
llama_perf_context_print(ctx);
fprintf(stderr, "\n");
llama_sampler_free(smpl);
llama_free(ctx);
llama_model_free(model);
free(prompt_tokens);
}