KBaba7/llama.cpp
0
1#include "arg.h"2#include "log.h"3#include "common.h"4#include "sampling.h"5#include "clip.h"6#include "llava.h"7#include "llama.h"8#include "ggml.h"9 10#include <algorithm>11#include <cstdio>12#include <cstdlib>13#include <cstring>14#include <vector>15#include <iostream> // TODO: remove me16 17struct llava_context {18 struct clip_ctx * ctx_clip = NULL;19 struct llama_context * ctx_llama = NULL;20 struct llama_model * model = NULL;21};22 23static void show_additional_info(int /*argc*/, char ** argv) {24 LOG("\nexample usage:\n\n%s -m <llava-v1.5-7b/ggml-model-q5_k.gguf> --mmproj <llava-v1.5-7b/mmproj-model-f16.gguf> --image <path/to/an/image.jpg> --image <path/to/another/image.jpg> [--temp 0.1] [-p \"describe the image in detail.\"]\n", argv[0]);25 LOG("\nnote: a lower temperature value like 0.1 is recommended for better quality.\n");26}27 28static struct llama_model * llava_init(common_params * params) {29 llama_backend_init();30 llama_numa_init(params->numa);31 32 llama_model_params model_params = common_model_params_to_llama(*params);33 34 llama_model * model = llama_model_load_from_file(params->model.c_str(), model_params);35 if (model == NULL) {36 LOG_ERR("%s: unable to load model\n" , __func__);37 return NULL;38 }39 return model;40}41 42static struct llava_context * llava_init_context(common_params * params, llama_model * model) {43 auto prompt = params->prompt;44 if (prompt.empty()) {45 prompt = "describe the image in detail.";46 }47 48 llama_context_params ctx_params = common_context_params_to_llama(*params);49 if (params->n_ctx < 2048) {50 // warn user here, "Image processing requires at least 2048 context, setting context to 2048"51 LOG_WRN("%s: Image processing requires at least 2048 context, setting context to 2048\n" , __func__);52 ctx_params.n_ctx = 2048;53 } else {54 ctx_params.n_ctx = params->n_ctx;55 }56 57 llama_context * ctx_llama = llama_init_from_model(model, ctx_params);58 59 if (ctx_llama == NULL) {60 LOG_ERR("%s: failed to create the llama_context\n" , __func__);61 return NULL;62 }63 64 auto * ctx_llava = (struct llava_context *)malloc(sizeof(llava_context));65 66 ctx_llava->ctx_llama = ctx_llama;67 ctx_llava->model = model;68 return ctx_llava;69}70 71static void llava_free(struct llava_context * ctx_llava) {72 if (ctx_llava->ctx_clip) {73 clip_free(ctx_llava->ctx_clip);74 ctx_llava->ctx_clip = NULL;75 }76 77 llama_free(ctx_llava->ctx_llama);78 llama_model_free(ctx_llava->model);79 llama_backend_free();80}81 82static struct clip_ctx * clip_init_context(common_params * params) {83 const char * clip_path = params->mmproj.c_str();84 85 auto prompt = params->prompt;86 if (prompt.empty()) {87 prompt = "describe the image in detail.";88 }89 auto * ctx_clip = clip_model_load(clip_path, /*verbosity=*/ 1);90 return ctx_clip;91}92 93static bool eval_tokens(struct llama_context * ctx_llama, std::vector<llama_token> tokens, int n_batch, int * n_past) {94 int N = (int) tokens.size();95 for (int i = 0; i < N; i += n_batch) {96 int n_eval = (int) tokens.size() - i;97 if (n_eval > n_batch) {98 n_eval = n_batch;99 }100 if (llama_decode(ctx_llama, llama_batch_get_one(&tokens[i], n_eval))) {101 LOG_ERR("%s : failed to eval. token %d/%d (batch size %d, n_past %d)\n", __func__, i, N, n_batch, *n_past);102 return false;103 }104 *n_past += n_eval;105 }106 return true;107}108 109static bool eval_id(struct llama_context * ctx_llama, int id, int * n_past) {110 std::vector<llama_token> tokens;111 tokens.push_back(id);112 return eval_tokens(ctx_llama, tokens, 1, n_past);113}114 115static bool eval_string(struct llama_context * ctx_llama, const char* str, int n_batch, int * n_past, bool add_bos){116 std::string str2 = str;117 std::vector<llama_token> embd_inp = common_tokenize(ctx_llama, str2, add_bos, true);118 return eval_tokens(ctx_llama, embd_inp, n_batch, n_past);119}120 121static void process_eval_image_embed(struct llava_context * ctx_llava, const struct llava_image_embed * embeds, int n_batch, int * n_past, int idx) {122 float * image_embed = (float *)malloc(clip_embd_nbytes(ctx_llava->ctx_clip));123 std::memcpy(image_embed, embeds->embed + idx * clip_n_patches(ctx_llava->ctx_clip) * clip_n_mmproj_embd(ctx_llava->ctx_clip), clip_embd_nbytes(ctx_llava->ctx_clip));124 125 auto * slice_embed = (llava_image_embed*)malloc(sizeof(llava_image_embed));126 slice_embed->embed = image_embed;127 slice_embed->n_image_pos = clip_n_patches(ctx_llava->ctx_clip);128 llava_eval_image_embed(ctx_llava->ctx_llama, slice_embed, n_batch, n_past);129 llava_image_embed_free(slice_embed);130}131 132static void process_image(struct llava_context * ctx_llava, struct llava_image_embed * embeds, common_params * params, int &n_past) {133 std::string system_prompt;134 int idx = 0;135 int num_image_embeds = embeds->n_image_pos / clip_n_patches(ctx_llava->ctx_clip);136 int has_minicpmv_projector = clip_is_minicpmv(ctx_llava->ctx_clip);137 if (has_minicpmv_projector == 2) {138 system_prompt = "<|begin_of_text|><|start_header_id|>user<|end_header_id|>\n\n";139 }140 else if (has_minicpmv_projector == 3) {141 system_prompt = "<|im_start|>user\n";142 }143 else if (has_minicpmv_projector == 4) {144 system_prompt = "<|im_start|>user\n";145 }146 LOG_INF("%s: image token past: %d\n", __func__, n_past);147 eval_string(ctx_llava->ctx_llama, (system_prompt+"<image>").c_str(), params->n_batch, &n_past, false);148 process_eval_image_embed(ctx_llava, embeds, params->n_batch, &n_past, idx++);149 eval_string(ctx_llava->ctx_llama, std::string("</image>").c_str(), params->n_batch, &n_past, false);150 if (num_image_embeds > 1) {151 size_t num_image_embeds_col = clip_uhd_num_image_embeds_col(ctx_llava->ctx_clip);152 eval_string(ctx_llava->ctx_llama, std::string("<slice>").c_str(), params->n_batch, &n_past, false);153 for (size_t i = 0; i < (num_image_embeds-1)/num_image_embeds_col; ++i) {154 for (size_t j = 0; j < num_image_embeds_col; ++j) {155 eval_string(ctx_llava->ctx_llama, std::string("<image>").c_str(), params->n_batch, &n_past, false);156 process_eval_image_embed(ctx_llava, embeds, params->n_batch, &n_past, idx++);157 eval_string(ctx_llava->ctx_llama, std::string("</image>").c_str(), params->n_batch, &n_past, false);158 if (j == num_image_embeds_col - 1) {159 eval_string(ctx_llava->ctx_llama, std::string("\n").c_str(), params->n_batch, &n_past, false);160 }161 }162 }163 eval_string(ctx_llava->ctx_llama, std::string("</slice>").c_str(), params->n_batch, &n_past, false);164 }165 LOG_INF("%s: image token past: %d\n", __func__, n_past);166}167 168static const char * sample(struct common_sampler * smpl,169 struct llama_context * ctx_llama,170 int * n_past) {171 const llama_token id = common_sampler_sample(smpl, ctx_llama, -1);172 common_sampler_accept(smpl, id, true);173 174 const llama_model * model = llama_get_model(ctx_llama);175 const llama_vocab * vocab = llama_model_get_vocab(model);176 177 static std::string ret;178 if (llama_vocab_is_eog(vocab, id)) {179 ret = "</s>";180 } else {181 ret = common_token_to_piece(ctx_llama, id);182 }183 eval_id(ctx_llama, id, n_past);184 return ret.c_str();185}186 187static struct llava_context * minicpmv_init(common_params * params, const std::string & fname, int &n_past){188 auto * ctx_clip = clip_init_context(params);189 auto * embeds = llava_image_embed_make_with_filename(ctx_clip, params->cpuparams.n_threads, fname.c_str());190 if (!embeds) {191 LOG_ERR("failed to load image %s. Terminating\n\n", fname.c_str());192 return NULL;193 }194 195 // process the prompt196 if (params->prompt.empty() && params->interactive == false) {197 LOG_ERR("prompt should be given or interactive mode should be on");198 return NULL;199 }200 201 auto * model = llava_init(params);202 if (model == NULL) {203 fprintf(stderr, "%s: error: failed to init minicpmv model\n", __func__);204 return NULL;205 }206 const int64_t t_llava_init_start_us = ggml_time_us();207 auto * ctx_llava = llava_init_context(params, model);208 ctx_llava->ctx_clip = ctx_clip;209 const int64_t t_llava_init_end_us = ggml_time_us();210 float t_llava_init_ms = (t_llava_init_end_us - t_llava_init_start_us) / 1000.0;211 LOG_INF("%s: llava init in %8.2f ms.\n", __func__, t_llava_init_ms);212 213 const int64_t t_process_image_start_us = ggml_time_us();214 process_image(ctx_llava, embeds, params, n_past);215 const int64_t t_process_image_end_us = ggml_time_us();216 float t_process_image_ms = (t_process_image_end_us - t_process_image_start_us) / 1000.0;217 LOG_INF("%s: llama process image in %8.2f ms.\n", __func__, t_process_image_ms);218 219 llava_image_embed_free(embeds);220 return ctx_llava;221}222 223static struct common_sampler * llama_init(struct llava_context * ctx_llava, common_params * params, const std::string & prompt, int & n_past, bool is_first = false){224 std::string user_prompt = prompt;225 int has_minicpmv_projector = clip_is_minicpmv(ctx_llava->ctx_clip);226 if (!is_first) {227 if (has_minicpmv_projector == 2) {228 user_prompt = "<|begin_of_text|><|start_header_id|>user<|end_header_id|>\n\n" + prompt;229 }230 else if (has_minicpmv_projector == 3) {231 user_prompt = "<|im_start|>user\n" + prompt;232 }233 else if (has_minicpmv_projector == 4) {234 user_prompt = "<|im_start|>user\n" + prompt;235 }236 }237 238 eval_string(ctx_llava->ctx_llama, user_prompt.c_str(), params->n_batch, &n_past, false);239 if (has_minicpmv_projector == 2) {240 eval_string(ctx_llava->ctx_llama, "<|eot_id|><|start_header_id|>assistant<|end_header_id|>\n\n", params->n_batch, &n_past, false);241 }242 else if (has_minicpmv_projector == 3) {243 eval_string(ctx_llava->ctx_llama, "<|im_end|><|im_start|>assistant\n", params->n_batch, &n_past, false);244 }245 else if (has_minicpmv_projector == 4) {246 eval_string(ctx_llava->ctx_llama, "<|im_end|><|im_start|>assistant\n", params->n_batch, &n_past, false);247 }248 249 // generate the response250 251 LOG_INF("\n");252 253 struct common_sampler * smpl = common_sampler_init(ctx_llava->model, params->sampling);254 return smpl;255}256 257static const char * llama_loop(struct llava_context * ctx_llava,struct common_sampler * smpl, int &n_past){258 259 const char * tmp = sample(smpl, ctx_llava->ctx_llama, &n_past);260 return tmp;261}262 263int main(int argc, char ** argv) {264 ggml_time_init();265 266 common_params params;267 268 if (!common_params_parse(argc, argv, params, LLAMA_EXAMPLE_LLAVA, show_additional_info)) {269 return 1;270 }271 272 common_init();273 274 if (params.mmproj.empty() || (params.image.empty())) {275 show_additional_info(argc, argv);276 return 1;277 }278 279 for (auto & image : params.image) {280 int n_past = 0;281 auto * ctx_llava = minicpmv_init(¶ms, image, n_past);282 283 if (!params.prompt.empty()) {284 LOG("<user>%s\n", params.prompt.c_str());285 LOG("<assistant>");286 auto * smpl = llama_init(ctx_llava, ¶ms, params.prompt, n_past, true);287 const int max_tgt_len = params.n_predict < 0 ? 256 : params.n_predict;288 std::string response;289 bool have_tmp = false;290 for (int i = 0; i < max_tgt_len; i++) {291 const auto * tmp = llama_loop(ctx_llava, smpl, n_past);292 response += tmp;293 if (strcmp(tmp, "</s>") == 0){294 if (!have_tmp) {295 continue;296 }297 break;298 }299 if (strstr(tmp, "###")) break; // Yi-VL behavior300 have_tmp = true;301 printf("%s", tmp);302 if (strstr(response.c_str(), "<user>")) break; // minicpm-v303 304 fflush(stdout);305 }306 common_sampler_free(smpl);307 }else {308 while (true) {309 LOG("<user>");310 std::string prompt;311 std::getline(std::cin, prompt);312 LOG("<assistant>");313 auto * smpl = llama_init(ctx_llava, ¶ms, prompt, n_past, true);314 const int max_tgt_len = params.n_predict < 0 ? 256 : params.n_predict;315 std::string response;316 for (int i = 0; i < max_tgt_len; i++) {317 const auto * tmp = llama_loop(ctx_llava, smpl, n_past);318 response += tmp;319 if (strcmp(tmp, "</s>") == 0) break;320 printf("%s", tmp);// mistral llava-1.6321 if (strstr(response.c_str(), "<user>")) break; // minicpm-v322 fflush(stdout);323 }324 common_sampler_free(smpl);325 }326 }327 printf("\n");328 llama_perf_context_print(ctx_llava->ctx_llama);329 330 ctx_llava->model = NULL;331 llava_free(ctx_llava);332 }333 334 return 0;335}336 