KBaba7/llama.cpp
0
1#if defined(_WIN32)2# include <windows.h>3# include <io.h>4#else5# include <sys/file.h>6# include <sys/ioctl.h>7# include <unistd.h>8#endif9 10#if defined(LLAMA_USE_CURL)11# include <curl/curl.h>12#endif13 14#include <signal.h>15 16#include <climits>17#include <cstdarg>18#include <cstdio>19#include <cstring>20#include <filesystem>21#include <iostream>22#include <list>23#include <sstream>24#include <string>25#include <vector>26 27#include "chat-template.hpp"28#include "common.h"29#include "json.hpp"30#include "linenoise.cpp/linenoise.h"31#include "llama-cpp.h"32#include "log.h"33 34#if defined(__unix__) || (defined(__APPLE__) && defined(__MACH__)) || defined(_WIN32)35[[noreturn]] static void sigint_handler(int) {36 printf("\n" LOG_COL_DEFAULT);37 exit(0); // not ideal, but it's the only way to guarantee exit in all cases38}39#endif40 41GGML_ATTRIBUTE_FORMAT(1, 2)42static std::string fmt(const char * fmt, ...) {43 va_list ap;44 va_list ap2;45 va_start(ap, fmt);46 va_copy(ap2, ap);47 const int size = vsnprintf(NULL, 0, fmt, ap);48 GGML_ASSERT(size >= 0 && size < INT_MAX); // NOLINT49 std::string buf;50 buf.resize(size);51 const int size2 = vsnprintf(const_cast<char *>(buf.data()), buf.size() + 1, fmt, ap2);52 GGML_ASSERT(size2 == size);53 va_end(ap2);54 va_end(ap);55 56 return buf;57}58 59GGML_ATTRIBUTE_FORMAT(1, 2)60static int printe(const char * fmt, ...) {61 va_list args;62 va_start(args, fmt);63 const int ret = vfprintf(stderr, fmt, args);64 va_end(args);65 66 return ret;67}68 69static std::string strftime_fmt(const char * fmt, const std::tm & tm) {70 std::ostringstream oss;71 oss << std::put_time(&tm, fmt);72 73 return oss.str();74}75 76class Opt {77 public:78 int init(int argc, const char ** argv) {79 ctx_params = llama_context_default_params();80 model_params = llama_model_default_params();81 context_size_default = ctx_params.n_batch;82 ngl_default = model_params.n_gpu_layers;83 common_params_sampling sampling;84 temperature_default = sampling.temp;85 86 if (argc < 2) {87 printe("Error: No arguments provided.\n");88 print_help();89 return 1;90 }91 92 // Parse arguments93 if (parse(argc, argv)) {94 printe("Error: Failed to parse arguments.\n");95 print_help();96 return 1;97 }98 99 // If help is requested, show help and exit100 if (help) {101 print_help();102 return 2;103 }104 105 ctx_params.n_batch = context_size >= 0 ? context_size : context_size_default;106 ctx_params.n_ctx = ctx_params.n_batch;107 model_params.n_gpu_layers = ngl >= 0 ? ngl : ngl_default;108 temperature = temperature >= 0 ? temperature : temperature_default;109 110 return 0; // Success111 }112 113 llama_context_params ctx_params;114 llama_model_params model_params;115 std::string model_;116 std::string user;117 bool use_jinja = false;118 int context_size = -1, ngl = -1;119 float temperature = -1;120 bool verbose = false;121 122 private:123 int context_size_default = -1, ngl_default = -1;124 float temperature_default = -1;125 bool help = false;126 127 bool parse_flag(const char ** argv, int i, const char * short_opt, const char * long_opt) {128 return strcmp(argv[i], short_opt) == 0 || strcmp(argv[i], long_opt) == 0;129 }130 131 int handle_option_with_value(int argc, const char ** argv, int & i, int & option_value) {132 if (i + 1 >= argc) {133 return 1;134 }135 136 option_value = std::atoi(argv[++i]);137 138 return 0;139 }140 141 int handle_option_with_value(int argc, const char ** argv, int & i, float & option_value) {142 if (i + 1 >= argc) {143 return 1;144 }145 146 option_value = std::atof(argv[++i]);147 148 return 0;149 }150 151 int parse(int argc, const char ** argv) {152 bool options_parsing = true;153 for (int i = 1, positional_args_i = 0; i < argc; ++i) {154 if (options_parsing && (strcmp(argv[i], "-c") == 0 || strcmp(argv[i], "--context-size") == 0)) {155 if (handle_option_with_value(argc, argv, i, context_size) == 1) {156 return 1;157 }158 } else if (options_parsing &&159 (strcmp(argv[i], "-n") == 0 || strcmp(argv[i], "-ngl") == 0 || strcmp(argv[i], "--ngl") == 0)) {160 if (handle_option_with_value(argc, argv, i, ngl) == 1) {161 return 1;162 }163 } else if (options_parsing && strcmp(argv[i], "--temp") == 0) {164 if (handle_option_with_value(argc, argv, i, temperature) == 1) {165 return 1;166 }167 } else if (options_parsing &&168 (parse_flag(argv, i, "-v", "--verbose") || parse_flag(argv, i, "-v", "--log-verbose"))) {169 verbose = true;170 } else if (options_parsing && strcmp(argv[i], "--jinja") == 0) {171 use_jinja = true;172 } else if (options_parsing && parse_flag(argv, i, "-h", "--help")) {173 help = true;174 return 0;175 } else if (options_parsing && strcmp(argv[i], "--") == 0) {176 options_parsing = false;177 } else if (positional_args_i == 0) {178 if (!argv[i][0] || argv[i][0] == '-') {179 return 1;180 }181 182 ++positional_args_i;183 model_ = argv[i];184 } else if (positional_args_i == 1) {185 ++positional_args_i;186 user = argv[i];187 } else {188 user += " " + std::string(argv[i]);189 }190 }191 192 if (model_.empty()){193 return 1;194 }195 196 return 0;197 }198 199 void print_help() const {200 printf(201 "Description:\n"202 " Runs a llm\n"203 "\n"204 "Usage:\n"205 " llama-run [options] model [prompt]\n"206 "\n"207 "Options:\n"208 " -c, --context-size <value>\n"209 " Context size (default: %d)\n"210 " -n, -ngl, --ngl <value>\n"211 " Number of GPU layers (default: %d)\n"212 " --temp <value>\n"213 " Temperature (default: %.1f)\n"214 " -v, --verbose, --log-verbose\n"215 " Set verbosity level to infinity (i.e. log all messages, useful for debugging)\n"216 " -h, --help\n"217 " Show help message\n"218 "\n"219 "Commands:\n"220 " model\n"221 " Model is a string with an optional prefix of \n"222 " huggingface:// (hf://), ollama://, https:// or file://.\n"223 " If no protocol is specified and a file exists in the specified\n"224 " path, file:// is assumed, otherwise if a file does not exist in\n"225 " the specified path, ollama:// is assumed. Models that are being\n"226 " pulled are downloaded with .partial extension while being\n"227 " downloaded and then renamed as the file without the .partial\n"228 " extension when complete.\n"229 "\n"230 "Examples:\n"231 " llama-run llama3\n"232 " llama-run ollama://granite-code\n"233 " llama-run ollama://smollm:135m\n"234 " llama-run hf://QuantFactory/SmolLM-135M-GGUF/SmolLM-135M.Q2_K.gguf\n"235 " llama-run "236 "huggingface://bartowski/SmolLM-1.7B-Instruct-v0.2-GGUF/SmolLM-1.7B-Instruct-v0.2-IQ3_M.gguf\n"237 " llama-run https://example.com/some-file1.gguf\n"238 " llama-run some-file2.gguf\n"239 " llama-run file://some-file3.gguf\n"240 " llama-run --ngl 999 some-file4.gguf\n"241 " llama-run --ngl 999 some-file5.gguf Hello World\n",242 context_size_default, ngl_default, temperature_default);243 }244};245 246struct progress_data {247 size_t file_size = 0;248 std::chrono::steady_clock::time_point start_time = std::chrono::steady_clock::now();249 bool printed = false;250};251 252static int get_terminal_width() {253#if defined(_WIN32)254 CONSOLE_SCREEN_BUFFER_INFO csbi;255 GetConsoleScreenBufferInfo(GetStdHandle(STD_OUTPUT_HANDLE), &csbi);256 return csbi.srWindow.Right - csbi.srWindow.Left + 1;257#else258 struct winsize w;259 ioctl(STDOUT_FILENO, TIOCGWINSZ, &w);260 return w.ws_col;261#endif262}263 264#ifdef LLAMA_USE_CURL265class File {266 public:267 FILE * file = nullptr;268 269 FILE * open(const std::string & filename, const char * mode) {270 file = fopen(filename.c_str(), mode);271 272 return file;273 }274 275 int lock() {276 if (file) {277# ifdef _WIN32278 fd = _fileno(file);279 hFile = (HANDLE) _get_osfhandle(fd);280 if (hFile == INVALID_HANDLE_VALUE) {281 fd = -1;282 283 return 1;284 }285 286 OVERLAPPED overlapped = {};287 if (!LockFileEx(hFile, LOCKFILE_EXCLUSIVE_LOCK | LOCKFILE_FAIL_IMMEDIATELY, 0, MAXDWORD, MAXDWORD,288 &overlapped)) {289 fd = -1;290 291 return 1;292 }293# else294 fd = fileno(file);295 if (flock(fd, LOCK_EX | LOCK_NB) != 0) {296 fd = -1;297 298 return 1;299 }300# endif301 }302 303 return 0;304 }305 306 ~File() {307 if (fd >= 0) {308# ifdef _WIN32309 if (hFile != INVALID_HANDLE_VALUE) {310 OVERLAPPED overlapped = {};311 UnlockFileEx(hFile, 0, MAXDWORD, MAXDWORD, &overlapped);312 }313# else314 flock(fd, LOCK_UN);315# endif316 }317 318 if (file) {319 fclose(file);320 }321 }322 323 private:324 int fd = -1;325# ifdef _WIN32326 HANDLE hFile = nullptr;327# endif328};329 330class HttpClient {331 public:332 int init(const std::string & url, const std::vector<std::string> & headers, const std::string & output_file,333 const bool progress, std::string * response_str = nullptr) {334 if (std::filesystem::exists(output_file)) {335 return 0;336 }337 338 std::string output_file_partial;339 curl = curl_easy_init();340 if (!curl) {341 return 1;342 }343 344 progress_data data;345 File out;346 if (!output_file.empty()) {347 output_file_partial = output_file + ".partial";348 if (!out.open(output_file_partial, "ab")) {349 printe("Failed to open file for writing\n");350 351 return 1;352 }353 354 if (out.lock()) {355 printe("Failed to exclusively lock file\n");356 357 return 1;358 }359 }360 361 set_write_options(response_str, out);362 data.file_size = set_resume_point(output_file_partial);363 set_progress_options(progress, data);364 set_headers(headers);365 CURLcode res = perform(url);366 if (res != CURLE_OK){367 printe("Fetching resource '%s' failed: %s\n", url.c_str(), curl_easy_strerror(res));368 return 1;369 }370 if (!output_file.empty()) {371 std::filesystem::rename(output_file_partial, output_file);372 }373 374 return 0;375 }376 377 ~HttpClient() {378 if (chunk) {379 curl_slist_free_all(chunk);380 }381 382 if (curl) {383 curl_easy_cleanup(curl);384 }385 }386 387 private:388 CURL * curl = nullptr;389 struct curl_slist * chunk = nullptr;390 391 void set_write_options(std::string * response_str, const File & out) {392 if (response_str) {393 curl_easy_setopt(curl, CURLOPT_WRITEFUNCTION, capture_data);394 curl_easy_setopt(curl, CURLOPT_WRITEDATA, response_str);395 } else {396 curl_easy_setopt(curl, CURLOPT_WRITEFUNCTION, write_data);397 curl_easy_setopt(curl, CURLOPT_WRITEDATA, out.file);398 }399 }400 401 size_t set_resume_point(const std::string & output_file) {402 size_t file_size = 0;403 if (std::filesystem::exists(output_file)) {404 file_size = std::filesystem::file_size(output_file);405 curl_easy_setopt(curl, CURLOPT_RESUME_FROM_LARGE, static_cast<curl_off_t>(file_size));406 }407 408 return file_size;409 }410 411 void set_progress_options(bool progress, progress_data & data) {412 if (progress) {413 curl_easy_setopt(curl, CURLOPT_NOPROGRESS, 0L);414 curl_easy_setopt(curl, CURLOPT_XFERINFODATA, &data);415 curl_easy_setopt(curl, CURLOPT_XFERINFOFUNCTION, update_progress);416 }417 }418 419 void set_headers(const std::vector<std::string> & headers) {420 if (!headers.empty()) {421 if (chunk) {422 curl_slist_free_all(chunk);423 chunk = 0;424 }425 426 for (const auto & header : headers) {427 chunk = curl_slist_append(chunk, header.c_str());428 }429 430 curl_easy_setopt(curl, CURLOPT_HTTPHEADER, chunk);431 }432 }433 434 CURLcode perform(const std::string & url) {435 curl_easy_setopt(curl, CURLOPT_URL, url.c_str());436 curl_easy_setopt(curl, CURLOPT_FOLLOWLOCATION, 1L);437 curl_easy_setopt(curl, CURLOPT_DEFAULT_PROTOCOL, "https");438 curl_easy_setopt(curl, CURLOPT_FAILONERROR, 1L);439 return curl_easy_perform(curl);440 }441 442 static std::string human_readable_time(double seconds) {443 int hrs = static_cast<int>(seconds) / 3600;444 int mins = (static_cast<int>(seconds) % 3600) / 60;445 int secs = static_cast<int>(seconds) % 60;446 447 if (hrs > 0) {448 return fmt("%dh %02dm %02ds", hrs, mins, secs);449 } else if (mins > 0) {450 return fmt("%dm %02ds", mins, secs);451 } else {452 return fmt("%ds", secs);453 }454 }455 456 static std::string human_readable_size(curl_off_t size) {457 static const char * suffix[] = { "B", "KB", "MB", "GB", "TB" };458 char length = sizeof(suffix) / sizeof(suffix[0]);459 int i = 0;460 double dbl_size = size;461 if (size > 1024) {462 for (i = 0; (size / 1024) > 0 && i < length - 1; i++, size /= 1024) {463 dbl_size = size / 1024.0;464 }465 }466 467 return fmt("%.2f %s", dbl_size, suffix[i]);468 }469 470 static int update_progress(void * ptr, curl_off_t total_to_download, curl_off_t now_downloaded, curl_off_t,471 curl_off_t) {472 progress_data * data = static_cast<progress_data *>(ptr);473 if (total_to_download <= 0) {474 return 0;475 }476 477 total_to_download += data->file_size;478 const curl_off_t now_downloaded_plus_file_size = now_downloaded + data->file_size;479 const curl_off_t percentage = calculate_percentage(now_downloaded_plus_file_size, total_to_download);480 std::string progress_prefix = generate_progress_prefix(percentage);481 482 const double speed = calculate_speed(now_downloaded, data->start_time);483 const double tim = (total_to_download - now_downloaded) / speed;484 std::string progress_suffix =485 generate_progress_suffix(now_downloaded_plus_file_size, total_to_download, speed, tim);486 487 int progress_bar_width = calculate_progress_bar_width(progress_prefix, progress_suffix);488 std::string progress_bar;489 generate_progress_bar(progress_bar_width, percentage, progress_bar);490 491 print_progress(progress_prefix, progress_bar, progress_suffix);492 data->printed = true;493 494 return 0;495 }496 497 static curl_off_t calculate_percentage(curl_off_t now_downloaded_plus_file_size, curl_off_t total_to_download) {498 return (now_downloaded_plus_file_size * 100) / total_to_download;499 }500 501 static std::string generate_progress_prefix(curl_off_t percentage) { return fmt("%3ld%% |", static_cast<long int>(percentage)); }502 503 static double calculate_speed(curl_off_t now_downloaded, const std::chrono::steady_clock::time_point & start_time) {504 const auto now = std::chrono::steady_clock::now();505 const std::chrono::duration<double> elapsed_seconds = now - start_time;506 return now_downloaded / elapsed_seconds.count();507 }508 509 static std::string generate_progress_suffix(curl_off_t now_downloaded_plus_file_size, curl_off_t total_to_download,510 double speed, double estimated_time) {511 const int width = 10;512 return fmt("%*s/%*s%*s/s%*s", width, human_readable_size(now_downloaded_plus_file_size).c_str(), width,513 human_readable_size(total_to_download).c_str(), width, human_readable_size(speed).c_str(), width,514 human_readable_time(estimated_time).c_str());515 }516 517 static int calculate_progress_bar_width(const std::string & progress_prefix, const std::string & progress_suffix) {518 int progress_bar_width = get_terminal_width() - progress_prefix.size() - progress_suffix.size() - 3;519 if (progress_bar_width < 1) {520 progress_bar_width = 1;521 }522 523 return progress_bar_width;524 }525 526 static std::string generate_progress_bar(int progress_bar_width, curl_off_t percentage,527 std::string & progress_bar) {528 const curl_off_t pos = (percentage * progress_bar_width) / 100;529 for (int i = 0; i < progress_bar_width; ++i) {530 progress_bar.append((i < pos) ? "โ" : " ");531 }532 533 return progress_bar;534 }535 536 static void print_progress(const std::string & progress_prefix, const std::string & progress_bar,537 const std::string & progress_suffix) {538 printe("\r%*s\r%s%s| %s", get_terminal_width(), " ", progress_prefix.c_str(), progress_bar.c_str(),539 progress_suffix.c_str());540 }541 // Function to write data to a file542 static size_t write_data(void * ptr, size_t size, size_t nmemb, void * stream) {543 FILE * out = static_cast<FILE *>(stream);544 return fwrite(ptr, size, nmemb, out);545 }546 547 // Function to capture data into a string548 static size_t capture_data(void * ptr, size_t size, size_t nmemb, void * stream) {549 std::string * str = static_cast<std::string *>(stream);550 str->append(static_cast<char *>(ptr), size * nmemb);551 return size * nmemb;552 }553};554#endif555 556class LlamaData {557 public:558 llama_model_ptr model;559 llama_sampler_ptr sampler;560 llama_context_ptr context;561 std::vector<llama_chat_message> messages;562 std::list<std::string> msg_strs;563 std::vector<char> fmtted;564 565 int init(Opt & opt) {566 model = initialize_model(opt);567 if (!model) {568 return 1;569 }570 571 context = initialize_context(model, opt);572 if (!context) {573 return 1;574 }575 576 sampler = initialize_sampler(opt);577 578 return 0;579 }580 581 private:582#ifdef LLAMA_USE_CURL583 int download(const std::string & url, const std::string & output_file, const bool progress,584 const std::vector<std::string> & headers = {}, std::string * response_str = nullptr) {585 HttpClient http;586 if (http.init(url, headers, output_file, progress, response_str)) {587 return 1;588 }589 590 return 0;591 }592#else593 int download(const std::string &, const std::string &, const bool, const std::vector<std::string> & = {},594 std::string * = nullptr) {595 printe("%s: llama.cpp built without libcurl, downloading from an url not supported.\n", __func__);596 597 return 1;598 }599#endif600 601 // Helper function to handle model tag extraction and URL construction602 std::pair<std::string, std::string> extract_model_and_tag(std::string & model, const std::string & base_url) {603 std::string model_tag = "latest";604 const size_t colon_pos = model.find(':');605 if (colon_pos != std::string::npos) {606 model_tag = model.substr(colon_pos + 1);607 model = model.substr(0, colon_pos);608 }609 610 std::string url = base_url + model + "/manifests/" + model_tag;611 612 return { model, url };613 }614 615 // Helper function to download and parse the manifest616 int download_and_parse_manifest(const std::string & url, const std::vector<std::string> & headers,617 nlohmann::json & manifest) {618 std::string manifest_str;619 int ret = download(url, "", false, headers, &manifest_str);620 if (ret) {621 return ret;622 }623 624 manifest = nlohmann::json::parse(manifest_str);625 626 return 0;627 }628 629 int huggingface_dl(std::string & model, const std::string & bn) {630 // Find the second occurrence of '/' after protocol string631 size_t pos = model.find('/');632 pos = model.find('/', pos + 1);633 std::string hfr, hff;634 std::vector<std::string> headers = { "User-Agent: llama-cpp", "Accept: application/json" };635 std::string url;636 637 if (pos == std::string::npos) {638 auto [model_name, manifest_url] = extract_model_and_tag(model, "https://huggingface.co/v2/");639 hfr = model_name;640 641 nlohmann::json manifest;642 int ret = download_and_parse_manifest(manifest_url, headers, manifest);643 if (ret) {644 return ret;645 }646 647 hff = manifest["ggufFile"]["rfilename"];648 } else {649 hfr = model.substr(0, pos);650 hff = model.substr(pos + 1);651 }652 653 url = "https://huggingface.co/" + hfr + "/resolve/main/" + hff;654 655 return download(url, bn, true, headers);656 }657 658 int ollama_dl(std::string & model, const std::string & bn) {659 const std::vector<std::string> headers = { "Accept: application/vnd.docker.distribution.manifest.v2+json" };660 if (model.find('/') == std::string::npos) {661 model = "library/" + model;662 }663 664 auto [model_name, manifest_url] = extract_model_and_tag(model, "https://registry.ollama.ai/v2/");665 nlohmann::json manifest;666 int ret = download_and_parse_manifest(manifest_url, {}, manifest);667 if (ret) {668 return ret;669 }670 671 std::string layer;672 for (const auto & l : manifest["layers"]) {673 if (l["mediaType"] == "application/vnd.ollama.image.model") {674 layer = l["digest"];675 break;676 }677 }678 679 std::string blob_url = "https://registry.ollama.ai/v2/" + model_name + "/blobs/" + layer;680 681 return download(blob_url, bn, true, headers);682 }683 684 int github_dl(const std::string & model, const std::string & bn) {685 std::string repository = model;686 std::string branch = "main";687 const size_t at_pos = model.find('@');688 if (at_pos != std::string::npos) {689 repository = model.substr(0, at_pos);690 branch = model.substr(at_pos + 1);691 }692 693 const std::vector<std::string> repo_parts = string_split(repository, "/");694 if (repo_parts.size() < 3) {695 printe("Invalid GitHub repository format\n");696 return 1;697 }698 699 const std::string & org = repo_parts[0];700 const std::string & project = repo_parts[1];701 std::string url = "https://raw.githubusercontent.com/" + org + "/" + project + "/" + branch;702 for (size_t i = 2; i < repo_parts.size(); ++i) {703 url += "/" + repo_parts[i];704 }705 706 return download(url, bn, true);707 }708 709 int s3_dl(const std::string & model, const std::string & bn) {710 const size_t slash_pos = model.find('/');711 if (slash_pos == std::string::npos) {712 return 1;713 }714 715 const std::string bucket = model.substr(0, slash_pos);716 const std::string key = model.substr(slash_pos + 1);717 const char * access_key = std::getenv("AWS_ACCESS_KEY_ID");718 const char * secret_key = std::getenv("AWS_SECRET_ACCESS_KEY");719 if (!access_key || !secret_key) {720 printe("AWS credentials not found in environment\n");721 return 1;722 }723 724 // Generate AWS Signature Version 4 headers725 // (Implementation requires HMAC-SHA256 and date handling)726 // Get current timestamp727 const time_t now = time(nullptr);728 const tm tm = *gmtime(&now);729 const std::string date = strftime_fmt("%Y%m%d", tm);730 const std::string datetime = strftime_fmt("%Y%m%dT%H%M%SZ", tm);731 const std::vector<std::string> headers = {732 "Authorization: AWS4-HMAC-SHA256 Credential=" + std::string(access_key) + "/" + date +733 "/us-east-1/s3/aws4_request",734 "x-amz-content-sha256: UNSIGNED-PAYLOAD", "x-amz-date: " + datetime735 };736 737 const std::string url = "https://" + bucket + ".s3.amazonaws.com/" + key;738 739 return download(url, bn, true, headers);740 }741 742 std::string basename(const std::string & path) {743 const size_t pos = path.find_last_of("/\\");744 if (pos == std::string::npos) {745 return path;746 }747 748 return path.substr(pos + 1);749 }750 751 int rm_until_substring(std::string & model_, const std::string & substring) {752 const std::string::size_type pos = model_.find(substring);753 if (pos == std::string::npos) {754 return 1;755 }756 757 model_ = model_.substr(pos + substring.size()); // Skip past the substring758 return 0;759 }760 761 int resolve_model(std::string & model_) {762 int ret = 0;763 if (string_starts_with(model_, "file://") || std::filesystem::exists(model_)) {764 rm_until_substring(model_, "://");765 766 return ret;767 }768 769 const std::string bn = basename(model_);770 if (string_starts_with(model_, "hf://") || string_starts_with(model_, "huggingface://") ||771 string_starts_with(model_, "hf.co/")) {772 rm_until_substring(model_, "hf.co/");773 rm_until_substring(model_, "://");774 ret = huggingface_dl(model_, bn);775 } else if ((string_starts_with(model_, "https://") || string_starts_with(model_, "http://")) &&776 !string_starts_with(model_, "https://ollama.com/library/")) {777 ret = download(model_, bn, true);778 } else if (string_starts_with(model_, "github:") || string_starts_with(model_, "github://")) {779 rm_until_substring(model_, "github:");780 rm_until_substring(model_, "://");781 ret = github_dl(model_, bn);782 } else if (string_starts_with(model_, "s3://")) {783 rm_until_substring(model_, "://");784 ret = s3_dl(model_, bn);785 } else { // ollama:// or nothing786 rm_until_substring(model_, "ollama.com/library/");787 rm_until_substring(model_, "://");788 ret = ollama_dl(model_, bn);789 }790 791 model_ = bn;792 793 return ret;794 }795 796 // Initializes the model and returns a unique pointer to it797 llama_model_ptr initialize_model(Opt & opt) {798 ggml_backend_load_all();799 resolve_model(opt.model_);800 printe(801 "\r%*s"802 "\rLoading model",803 get_terminal_width(), " ");804 llama_model_ptr model(llama_model_load_from_file(opt.model_.c_str(), opt.model_params));805 if (!model) {806 printe("%s: error: unable to load model from file: %s\n", __func__, opt.model_.c_str());807 }808 809 printe("\r%*s\r", static_cast<int>(sizeof("Loading model")), " ");810 return model;811 }812 813 // Initializes the context with the specified parameters814 llama_context_ptr initialize_context(const llama_model_ptr & model, const Opt & opt) {815 llama_context_ptr context(llama_init_from_model(model.get(), opt.ctx_params));816 if (!context) {817 printe("%s: error: failed to create the llama_context\n", __func__);818 }819 820 return context;821 }822 823 // Initializes and configures the sampler824 llama_sampler_ptr initialize_sampler(const Opt & opt) {825 llama_sampler_ptr sampler(llama_sampler_chain_init(llama_sampler_chain_default_params()));826 llama_sampler_chain_add(sampler.get(), llama_sampler_init_min_p(0.05f, 1));827 llama_sampler_chain_add(sampler.get(), llama_sampler_init_temp(opt.temperature));828 llama_sampler_chain_add(sampler.get(), llama_sampler_init_dist(LLAMA_DEFAULT_SEED));829 830 return sampler;831 }832};833 834// Add a message to `messages` and store its content in `msg_strs`835static void add_message(const char * role, const std::string & text, LlamaData & llama_data) {836 llama_data.msg_strs.push_back(std::move(text));837 llama_data.messages.push_back({ role, llama_data.msg_strs.back().c_str() });838}839 840// Function to apply the chat template and resize `formatted` if needed841static int apply_chat_template(const common_chat_template & tmpl, LlamaData & llama_data, const bool append, bool use_jinja) {842 if (use_jinja) {843 json messages = json::array();844 for (const auto & msg : llama_data.messages) {845 messages.push_back({846 {"role", msg.role},847 {"content", msg.content},848 });849 }850 try {851 minja::chat_template_inputs tmpl_inputs;852 tmpl_inputs.messages = messages;853 tmpl_inputs.add_generation_prompt = append;854 855 minja::chat_template_options tmpl_opts;856 tmpl_opts.use_bos_token = false;857 tmpl_opts.use_eos_token = false;858 859 auto result = tmpl.apply(tmpl_inputs, tmpl_opts);860 llama_data.fmtted.resize(result.size() + 1);861 memcpy(llama_data.fmtted.data(), result.c_str(), result.size() + 1);862 return result.size();863 } catch (const std::exception & e) {864 printe("failed to render the chat template: %s\n", e.what());865 return -1;866 }867 }868 int result = llama_chat_apply_template(869 tmpl.source().c_str(), llama_data.messages.data(), llama_data.messages.size(), append,870 append ? llama_data.fmtted.data() : nullptr, append ? llama_data.fmtted.size() : 0);871 if (append && result > static_cast<int>(llama_data.fmtted.size())) {872 llama_data.fmtted.resize(result);873 result = llama_chat_apply_template(tmpl.source().c_str(), llama_data.messages.data(),874 llama_data.messages.size(), append, llama_data.fmtted.data(),875 llama_data.fmtted.size());876 }877 878 return result;879}880 881// Function to tokenize the prompt882static int tokenize_prompt(const llama_vocab * vocab, const std::string & prompt,883 std::vector<llama_token> & prompt_tokens, const LlamaData & llama_data) {884 const bool is_first = llama_get_kv_cache_used_cells(llama_data.context.get()) == 0;885 886 const int n_prompt_tokens = -llama_tokenize(vocab, prompt.c_str(), prompt.size(), NULL, 0, is_first, true);887 prompt_tokens.resize(n_prompt_tokens);888 if (llama_tokenize(vocab, prompt.c_str(), prompt.size(), prompt_tokens.data(), prompt_tokens.size(), is_first,889 true) < 0) {890 printe("failed to tokenize the prompt\n");891 return -1;892 }893 894 return n_prompt_tokens;895}896 897// Check if we have enough space in the context to evaluate this batch898static int check_context_size(const llama_context_ptr & ctx, const llama_batch & batch) {899 const int n_ctx = llama_n_ctx(ctx.get());900 const int n_ctx_used = llama_get_kv_cache_used_cells(ctx.get());901 if (n_ctx_used + batch.n_tokens > n_ctx) {902 printf(LOG_COL_DEFAULT "\n");903 printe("context size exceeded\n");904 return 1;905 }906 907 return 0;908}909 910// convert the token to a string911static int convert_token_to_string(const llama_vocab * vocab, const llama_token token_id, std::string & piece) {912 char buf[256];913 int n = llama_token_to_piece(vocab, token_id, buf, sizeof(buf), 0, true);914 if (n < 0) {915 printe("failed to convert token to piece\n");916 return 1;917 }918 919 piece = std::string(buf, n);920 return 0;921}922 923static void print_word_and_concatenate_to_response(const std::string & piece, std::string & response) {924 printf("%s", piece.c_str());925 fflush(stdout);926 response += piece;927}928 929// helper function to evaluate a prompt and generate a response930static int generate(LlamaData & llama_data, const std::string & prompt, std::string & response) {931 const llama_vocab * vocab = llama_model_get_vocab(llama_data.model.get());932 933 std::vector<llama_token> tokens;934 if (tokenize_prompt(vocab, prompt, tokens, llama_data) < 0) {935 return 1;936 }937 938 // prepare a batch for the prompt939 llama_batch batch = llama_batch_get_one(tokens.data(), tokens.size());940 llama_token new_token_id;941 while (true) {942 check_context_size(llama_data.context, batch);943 if (llama_decode(llama_data.context.get(), batch)) {944 printe("failed to decode\n");945 return 1;946 }947 948 // sample the next token, check is it an end of generation?949 new_token_id = llama_sampler_sample(llama_data.sampler.get(), llama_data.context.get(), -1);950 if (llama_vocab_is_eog(vocab, new_token_id)) {951 break;952 }953 954 std::string piece;955 if (convert_token_to_string(vocab, new_token_id, piece)) {956 return 1;957 }958 959 print_word_and_concatenate_to_response(piece, response);960 961 // prepare the next batch with the sampled token962 batch = llama_batch_get_one(&new_token_id, 1);963 }964 965 printf(LOG_COL_DEFAULT);966 return 0;967}968 969static int read_user_input(std::string & user_input) {970 static const char * prompt_prefix = "> ";971#ifdef WIN32972 printf(973 "\r%*s"974 "\r" LOG_COL_DEFAULT "%s",975 get_terminal_width(), " ", prompt_prefix);976 977 std::getline(std::cin, user_input);978 if (std::cin.eof()) {979 printf("\n");980 return 1;981 }982#else983 std::unique_ptr<char, decltype(&std::free)> line(const_cast<char *>(linenoise(prompt_prefix)), free);984 if (!line) {985 return 1;986 }987 988 user_input = line.get();989#endif990 991 if (user_input == "/bye") {992 return 1;993 }994 995 if (user_input.empty()) {996 return 2;997 }998 999#ifndef WIN321000 linenoiseHistoryAdd(line.get());1001#endif1002 1003 return 0; // Should have data in happy path1004}1005 1006// Function to generate a response based on the prompt1007static int generate_response(LlamaData & llama_data, const std::string & prompt, std::string & response,1008 const bool stdout_a_terminal) {1009 // Set response color1010 if (stdout_a_terminal) {1011 printf(LOG_COL_YELLOW);1012 }1013 1014 if (generate(llama_data, prompt, response)) {1015 printe("failed to generate response\n");1016 return 1;1017 }1018 1019 // End response with color reset and newline1020 printf("\n%s", stdout_a_terminal ? LOG_COL_DEFAULT : "");1021 return 0;1022}1023 1024// Helper function to apply the chat template and handle errors1025static int apply_chat_template_with_error_handling(const common_chat_template & tmpl, LlamaData & llama_data, const bool append, int & output_length, bool use_jinja) {1026 const int new_len = apply_chat_template(tmpl, llama_data, append, use_jinja);1027 if (new_len < 0) {1028 printe("failed to apply the chat template\n");1029 return -1;1030 }1031 1032 output_length = new_len;1033 return 0;1034}1035 1036// Helper function to handle user input1037static int handle_user_input(std::string & user_input, const std::string & user) {1038 if (!user.empty()) {1039 user_input = user;1040 return 0; // No need for interactive input1041 }1042 1043 return read_user_input(user_input); // Returns true if input ends the loop1044}1045 1046static bool is_stdin_a_terminal() {1047#if defined(_WIN32)1048 HANDLE hStdin = GetStdHandle(STD_INPUT_HANDLE);1049 DWORD mode;1050 return GetConsoleMode(hStdin, &mode);1051#else1052 return isatty(STDIN_FILENO);1053#endif1054}1055 1056static bool is_stdout_a_terminal() {1057#if defined(_WIN32)1058 HANDLE hStdout = GetStdHandle(STD_OUTPUT_HANDLE);1059 DWORD mode;1060 return GetConsoleMode(hStdout, &mode);1061#else1062 return isatty(STDOUT_FILENO);1063#endif1064}1065 1066// Function to handle user input1067static int get_user_input(std::string & user_input, const std::string & user) {1068 while (true) {1069 const int ret = handle_user_input(user_input, user);1070 if (ret == 1) {1071 return 1;1072 }1073 1074 if (ret == 2) {1075 continue;1076 }1077 1078 break;1079 }1080 1081 return 0;1082}1083 1084// Main chat loop function1085static int chat_loop(LlamaData & llama_data, const std::string & user, bool use_jinja) {1086 int prev_len = 0;1087 llama_data.fmtted.resize(llama_n_ctx(llama_data.context.get()));1088 auto chat_templates = common_chat_templates_from_model(llama_data.model.get(), "");1089 GGML_ASSERT(chat_templates.template_default);1090 static const bool stdout_a_terminal = is_stdout_a_terminal();1091 while (true) {1092 // Get user input1093 std::string user_input;1094 if (get_user_input(user_input, user) == 1) {1095 return 0;1096 }1097 1098 add_message("user", user.empty() ? user_input : user, llama_data);1099 int new_len;1100 if (apply_chat_template_with_error_handling(*chat_templates.template_default, llama_data, true, new_len, use_jinja) < 0) {1101 return 1;1102 }1103 1104 std::string prompt(llama_data.fmtted.begin() + prev_len, llama_data.fmtted.begin() + new_len);1105 std::string response;1106 if (generate_response(llama_data, prompt, response, stdout_a_terminal)) {1107 return 1;1108 }1109 1110 if (!user.empty()) {1111 break;1112 }1113 1114 add_message("assistant", response, llama_data);1115 if (apply_chat_template_with_error_handling(*chat_templates.template_default, llama_data, false, prev_len, use_jinja) < 0) {1116 return 1;1117 }1118 }1119 1120 return 0;1121}1122 1123static void log_callback(const enum ggml_log_level level, const char * text, void * p) {1124 const Opt * opt = static_cast<Opt *>(p);1125 if (opt->verbose || level == GGML_LOG_LEVEL_ERROR) {1126 printe("%s", text);1127 }1128}1129 1130static std::string read_pipe_data() {1131 std::ostringstream result;1132 result << std::cin.rdbuf(); // Read all data from std::cin1133 return result.str();1134}1135 1136static void ctrl_c_handling() {1137#if defined(__unix__) || (defined(__APPLE__) && defined(__MACH__))1138 struct sigaction sigint_action;1139 sigint_action.sa_handler = sigint_handler;1140 sigemptyset(&sigint_action.sa_mask);1141 sigint_action.sa_flags = 0;1142 sigaction(SIGINT, &sigint_action, NULL);1143#elif defined(_WIN32)1144 auto console_ctrl_handler = +[](DWORD ctrl_type) -> BOOL {1145 return (ctrl_type == CTRL_C_EVENT) ? (sigint_handler(SIGINT), true) : false;1146 };1147 SetConsoleCtrlHandler(reinterpret_cast<PHANDLER_ROUTINE>(console_ctrl_handler), true);1148#endif1149}1150 1151int main(int argc, const char ** argv) {1152 ctrl_c_handling();1153 Opt opt;1154 const int ret = opt.init(argc, argv);1155 if (ret == 2) {1156 return 0;1157 } else if (ret) {1158 return 1;1159 }1160 1161 if (!is_stdin_a_terminal()) {1162 if (!opt.user.empty()) {1163 opt.user += "\n\n";1164 }1165 1166 opt.user += read_pipe_data();1167 }1168 1169 llama_log_set(log_callback, &opt);1170 LlamaData llama_data;1171 if (llama_data.init(opt)) {1172 return 1;1173 }1174 1175 if (chat_loop(llama_data, opt.user, opt.use_jinja)) {1176 return 1;1177 }1178 1179 return 0;1180}1181 