Team Ai
Apppublic

KBaba7/llama.cpp

sourceHugging Faceapache-2.0updated 2y agoView on Hugging Face
0likes
simple.cpp207 linesDownload Raw Back to simple
1#include "llama.h"2#include <cstdio>3#include <cstring>4#include <string>5#include <vector>6 7static void print_usage(int, char ** argv) {8    printf("\nexample usage:\n");9    printf("\n    %s -m model.gguf [-n n_predict] [-ngl n_gpu_layers] [prompt]\n", argv[0]);10    printf("\n");11}12 13int main(int argc, char ** argv) {14    // path to the model gguf file15    std::string model_path;16    // prompt to generate text from17    std::string prompt = "Hello my name is";18    // number of layers to offload to the GPU19    int ngl = 99;20    // number of tokens to predict21    int n_predict = 32;22 23    // parse command line arguments24 25    {26        int i = 1;27        for (; i < argc; i++) {28            if (strcmp(argv[i], "-m") == 0) {29                if (i + 1 < argc) {30                    model_path = argv[++i];31                } else {32                    print_usage(argc, argv);33                    return 1;34                }35            } else if (strcmp(argv[i], "-n") == 0) {36                if (i + 1 < argc) {37                    try {38                        n_predict = std::stoi(argv[++i]);39                    } catch (...) {40                        print_usage(argc, argv);41                        return 1;42                    }43                } else {44                    print_usage(argc, argv);45                    return 1;46                }47            } else if (strcmp(argv[i], "-ngl") == 0) {48                if (i + 1 < argc) {49                    try {50                        ngl = std::stoi(argv[++i]);51                    } catch (...) {52                        print_usage(argc, argv);53                        return 1;54                    }55                } else {56                    print_usage(argc, argv);57                    return 1;58                }59            } else {60                // prompt starts here61                break;62            }63        }64        if (model_path.empty()) {65            print_usage(argc, argv);66            return 1;67        }68        if (i < argc) {69            prompt = argv[i++];70            for (; i < argc; i++) {71                prompt += " ";72                prompt += argv[i];73            }74        }75    }76 77    // load dynamic backends78 79    ggml_backend_load_all();80 81    // initialize the model82 83    llama_model_params model_params = llama_model_default_params();84    model_params.n_gpu_layers = ngl;85 86    llama_model * model = llama_model_load_from_file(model_path.c_str(), model_params);87    const llama_vocab * vocab = llama_model_get_vocab(model);88 89    if (model == NULL) {90        fprintf(stderr , "%s: error: unable to load model\n" , __func__);91        return 1;92    }93 94    // tokenize the prompt95 96    // find the number of tokens in the prompt97    const int n_prompt = -llama_tokenize(vocab, prompt.c_str(), prompt.size(), NULL, 0, true, true);98 99    // allocate space for the tokens and tokenize the prompt100    std::vector<llama_token> prompt_tokens(n_prompt);101    if (llama_tokenize(vocab, prompt.c_str(), prompt.size(), prompt_tokens.data(), prompt_tokens.size(), true, true) < 0) {102        fprintf(stderr, "%s: error: failed to tokenize the prompt\n", __func__);103        return 1;104    }105 106    // initialize the context107 108    llama_context_params ctx_params = llama_context_default_params();109    // n_ctx is the context size110    ctx_params.n_ctx = n_prompt + n_predict - 1;111    // n_batch is the maximum number of tokens that can be processed in a single call to llama_decode112    ctx_params.n_batch = n_prompt;113    // enable performance counters114    ctx_params.no_perf = false;115 116    llama_context * ctx = llama_init_from_model(model, ctx_params);117 118    if (ctx == NULL) {119        fprintf(stderr , "%s: error: failed to create the llama_context\n" , __func__);120        return 1;121    }122 123    // initialize the sampler124 125    auto sparams = llama_sampler_chain_default_params();126    sparams.no_perf = false;127    llama_sampler * smpl = llama_sampler_chain_init(sparams);128 129    llama_sampler_chain_add(smpl, llama_sampler_init_greedy());130 131    // print the prompt token-by-token132 133    for (auto id : prompt_tokens) {134        char buf[128];135        int n = llama_token_to_piece(vocab, id, buf, sizeof(buf), 0, true);136        if (n < 0) {137            fprintf(stderr, "%s: error: failed to convert token to piece\n", __func__);138            return 1;139        }140        std::string s(buf, n);141        printf("%s", s.c_str());142    }143 144    // prepare a batch for the prompt145 146    llama_batch batch = llama_batch_get_one(prompt_tokens.data(), prompt_tokens.size());147 148    // main loop149 150    const auto t_main_start = ggml_time_us();151    int n_decode = 0;152    llama_token new_token_id;153 154    for (int n_pos = 0; n_pos + batch.n_tokens < n_prompt + n_predict; ) {155        // evaluate the current batch with the transformer model156        if (llama_decode(ctx, batch)) {157            fprintf(stderr, "%s : failed to eval, return code %d\n", __func__, 1);158            return 1;159        }160 161        n_pos += batch.n_tokens;162 163        // sample the next token164        {165            new_token_id = llama_sampler_sample(smpl, ctx, -1);166 167            // is it an end of generation?168            if (llama_vocab_is_eog(vocab, new_token_id)) {169                break;170            }171 172            char buf[128];173            int n = llama_token_to_piece(vocab, new_token_id, buf, sizeof(buf), 0, true);174            if (n < 0) {175                fprintf(stderr, "%s: error: failed to convert token to piece\n", __func__);176                return 1;177            }178            std::string s(buf, n);179            printf("%s", s.c_str());180            fflush(stdout);181 182            // prepare the next batch with the sampled token183            batch = llama_batch_get_one(&new_token_id, 1);184 185            n_decode += 1;186        }187    }188 189    printf("\n");190 191    const auto t_main_end = ggml_time_us();192 193    fprintf(stderr, "%s: decoded %d tokens in %.2f s, speed: %.2f t/s\n",194            __func__, n_decode, (t_main_end - t_main_start) / 1000000.0f, n_decode / ((t_main_end - t_main_start) / 1000000.0f));195 196    fprintf(stderr, "\n");197    llama_perf_sampler_print(smpl);198    llama_perf_context_print(ctx);199    fprintf(stderr, "\n");200 201    llama_sampler_free(smpl);202    llama_free(ctx);203    llama_model_free(model);204 205    return 0;206}207