KBaba7/llama.cpp
0
1#include "arg.h"2#include "common.h"3#include "llama.h"4 5#include <vector>6#include <cstdio>7 8int main(int argc, char ** argv) {9 common_params params;10 11 params.prompt = "The quick brown fox";12 params.sampling.seed = 1234;13 14 if (!common_params_parse(argc, argv, params, LLAMA_EXAMPLE_COMMON)) {15 return 1;16 }17 18 print_build_info();19 20 if (params.n_predict < 0) {21 params.n_predict = 16;22 }23 24 auto n_past = 0;25 26 std::string result0;27 std::string result1;28 std::string result2;29 30 // init31 common_init_result llama_init = common_init_from_params(params);32 33 llama_model * model = llama_init.model.get();34 llama_context * ctx = llama_init.context.get();35 36 if (model == nullptr || ctx == nullptr) {37 fprintf(stderr, "%s : failed to init\n", __func__);38 return 1;39 }40 41 auto sparams = llama_sampler_chain_default_params();42 43 llama_sampler * smpl = llama_sampler_chain_init(sparams);44 45 llama_sampler_chain_add(smpl, llama_sampler_init_dist(params.sampling.seed));46 47 // tokenize prompt48 auto tokens = common_tokenize(ctx, params.prompt, true);49 50 // prepare the batch51 llama_batch batch = llama_batch_init(tokens.size(), 0, 1);52 for (size_t i = 0; i < tokens.size(); i++) {53 common_batch_add(batch, tokens[i], i, {0}, false);54 }55 batch.logits[batch.n_tokens - 1] = true; // generate next token56 57 // evaluate prompt58 llama_decode(ctx, batch);59 n_past += batch.n_tokens;60 61 // save state (rng, logits, embedding and kv_cache) to file62 {63 std::vector<uint8_t> state_mem(llama_state_get_size(ctx));64 const size_t written = llama_state_get_data(ctx, state_mem.data(), state_mem.size());65 66 FILE *fp_write = fopen("dump_state.bin", "wb");67 fwrite(state_mem.data(), 1, written, fp_write);68 fclose(fp_write);69 70 fprintf(stderr, "%s : serialized state into %zd out of a maximum of %zd bytes\n", __func__, written, state_mem.size());71 }72 73 // save state (last tokens)74 const auto n_past_saved = n_past;75 76 // first run77 printf("\nfirst run: %s", params.prompt.c_str());78 79 for (auto i = 0; i < params.n_predict; i++) {80 auto next_token = llama_sampler_sample(smpl, ctx, -1);81 auto next_token_str = common_token_to_piece(ctx, next_token);82 83 printf("%s", next_token_str.c_str());84 result0 += next_token_str;85 86 common_batch_clear(batch);87 common_batch_add(batch, next_token, n_past, {0}, true);88 89 if (llama_decode(ctx, batch)) {90 fprintf(stderr, "\n%s : failed to evaluate\n", __func__);91 llama_batch_free(batch);92 return 1;93 }94 n_past += 1;95 }96 97 printf("\n\n");98 99 // make new context100 llama_context * ctx2 = llama_init_from_model(model, common_context_params_to_llama(params));101 102 llama_sampler * smpl2 = llama_sampler_chain_init(sparams);103 104 llama_sampler_chain_add(smpl2, llama_sampler_init_dist(params.sampling.seed));105 106 printf("\nsecond run: %s", params.prompt.c_str());107 108 // load state (rng, logits, embedding and kv_cache) from file109 {110 std::vector<uint8_t> state_mem;111 112 FILE * fp_read = fopen("dump_state.bin", "rb");113 fseek(fp_read, 0, SEEK_END);114 state_mem.resize(ftell(fp_read));115 fseek(fp_read, 0, SEEK_SET);116 const size_t read = fread(state_mem.data(), 1, state_mem.size(), fp_read);117 fclose(fp_read);118 119 if (read != llama_state_set_data(ctx2, state_mem.data(), state_mem.size())) {120 fprintf(stderr, "\n%s : failed to read state\n", __func__);121 return 1;122 }123 124 fprintf(stderr, "%s : deserialized state from %zd out of a maximum of %zd bytes\n", __func__, read, state_mem.size());125 }126 127 // restore state (last tokens)128 n_past = n_past_saved;129 130 // second run131 for (auto i = 0; i < params.n_predict; i++) {132 auto next_token = llama_sampler_sample(smpl2, ctx2, -1);133 auto next_token_str = common_token_to_piece(ctx2, next_token);134 135 printf("%s", next_token_str.c_str());136 result1 += next_token_str;137 138 common_batch_clear(batch);139 common_batch_add(batch, next_token, n_past, {0}, true);140 141 if (llama_decode(ctx2, batch)) {142 fprintf(stderr, "\n%s : failed to evaluate\n", __func__);143 llama_batch_free(batch);144 return 1;145 }146 n_past += 1;147 }148 149 printf("\n\n");150 151 if (result0 != result1) {152 fprintf(stderr, "\n%s : error : the 2 generations are different\n", __func__);153 return 1;154 }155 156 // make new context157 llama_context * ctx3 = llama_init_from_model(model, common_context_params_to_llama(params));158 159 llama_sampler * smpl3 = llama_sampler_chain_init(sparams);160 161 llama_sampler_chain_add(smpl3, llama_sampler_init_dist(params.sampling.seed));162 163 printf("\nsingle seq run: %s", params.prompt.c_str());164 165 // load state (rng, logits, embedding and kv_cache) from file166 {167 std::vector<uint8_t> state_mem;168 169 FILE * fp_read = fopen("dump_state.bin", "rb");170 fseek(fp_read, 0, SEEK_END);171 state_mem.resize(ftell(fp_read));172 fseek(fp_read, 0, SEEK_SET);173 const size_t read = fread(state_mem.data(), 1, state_mem.size(), fp_read);174 fclose(fp_read);175 176 if (read != llama_state_set_data(ctx3, state_mem.data(), state_mem.size())) {177 fprintf(stderr, "\n%s : failed to read state\n", __func__);178 return 1;179 }180 181 fprintf(stderr, "%s : deserialized state from %zd out of a maximum of %zd bytes\n", __func__, read, state_mem.size());182 }183 184 // restore state (last tokens)185 n_past = n_past_saved;186 187 // save seq 0 and load into seq 1188 {189 // save kv of seq 0190 std::vector<uint8_t> seq_store(llama_state_seq_get_size(ctx3, 0));191 const size_t ncopy = llama_state_seq_get_data(ctx3, seq_store.data(), seq_store.size(), 0);192 if (ncopy != seq_store.size()) {193 fprintf(stderr, "\n%s : seq copy data length %zd does not match expected length %zd\n", __func__, ncopy, seq_store.size());194 return 1;195 }196 fprintf(stderr, "%s : seq 0 copied, %zd bytes\n", __func__, ncopy);197 198 // erase whole kv199 llama_kv_cache_clear(ctx3);200 fprintf(stderr, "%s : kv cache cleared\n", __func__);201 202 // restore kv into seq 1203 const size_t nset = llama_state_seq_set_data(ctx3, seq_store.data(), seq_store.size(), 1);204 if (nset != seq_store.size()) {205 fprintf(stderr, "\n%s : seq set data length %zd does not match expected length %zd\n", __func__, nset, seq_store.size());206 return 1;207 }208 fprintf(stderr, "%s : seq 1 restored, %zd bytes\n", __func__, nset);209 }210 211 // third run with seq 1 instead of 0212 for (auto i = 0; i < params.n_predict; i++) {213 auto next_token = llama_sampler_sample(smpl3, ctx3, -1);214 auto next_token_str = common_token_to_piece(ctx3, next_token);215 216 printf("%s", next_token_str.c_str());217 result2 += next_token_str;218 219 common_batch_clear(batch);220 common_batch_add(batch, next_token, n_past, {1}, true);221 222 if (llama_decode(ctx3, batch)) {223 fprintf(stderr, "\n%s : failed to evaluate\n", __func__);224 llama_batch_free(batch);225 return 1;226 }227 n_past += 1;228 }229 230 printf("\n");231 232 llama_sampler_free(smpl);233 llama_sampler_free(smpl2);234 llama_sampler_free(smpl3);235 236 llama_batch_free(batch);237 238 if (result0 != result2) {239 fprintf(stderr, "\n%s : error : the seq restore generation is different\n", __func__);240 return 1;241 }242 243 fprintf(stderr, "\n%s : success\n", __func__);244 245 return 0;246}247 