Team Ai
Apppublic

KBaba7/llama.cpp

sourceHugging Faceapache-2.0updated 2y agoView on Hugging Face
0likes
save-load-state.cpp247 linesDownload Raw Back to save-load-state
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