diff --git a/src/llama-context.cpp b/src/llama-context.cpp index 52f8d53672a3..3b2b7be9d0f3 100644 --- a/src/llama-context.cpp +++ b/src/llama-context.cpp @@ -2999,28 +2999,31 @@ size_t llama_context::state_seq_get_data(llama_seq_id seq_id, uint8_t * dst, siz } size_t llama_context::state_seq_set_data(llama_seq_id seq_id, const uint8_t * src, size_t size, llama_state_seq_flags flags) { - std::unique_ptr io; - if (flags & LLAMA_STATE_SEQ_FLAGS_ON_DEVICE) { - // create a temporary io to read the magic and the src seq_id - io = std::make_unique(src, size); + try { + std::unique_ptr io; + if (flags & LLAMA_STATE_SEQ_FLAGS_ON_DEVICE) { + // create a temporary io to read the magic and the src seq_id + io = std::make_unique(src, size); + + uint32_t magic_read; + io->read(&magic_read, sizeof(magic_read)); + if (io_magic != magic_read) { + throw std::runtime_error("wrong sequence state magic"); + } - uint32_t magic_read; - io->read(&magic_read, sizeof(magic_read)); - if (io_magic != magic_read) { - throw std::runtime_error("wrong sequence state magic"); - } + llama_seq_id seq_id_read; + io->read(&seq_id_read, sizeof(seq_id_read)); - llama_seq_id seq_id_read; - io->read(&seq_id_read, sizeof(seq_id_read)); - - GGML_ASSERT(mem_storage.find(seq_id_read) != mem_storage.end()); + auto it = mem_storage.find(seq_id_read); + if (it == mem_storage.end()) { + throw std::runtime_error("sequence state not found"); + } - io = std::make_unique(src, size, mem_storage[seq_id_read]); - } else { - io = std::make_unique(src, size); - } + io = std::make_unique(src, size, it->second); + } else { + io = std::make_unique(src, size); + } - try { uint32_t magic_read; io->read(&magic_read, sizeof(magic_read)); if (io_magic != magic_read) { diff --git a/tests/CMakeLists.txt b/tests/CMakeLists.txt index e517d2c6359c..6cfd7e7e72d8 100644 --- a/tests/CMakeLists.txt +++ b/tests/CMakeLists.txt @@ -149,6 +149,7 @@ if (LLAMA_LLGUIDANCE) endif () llama_build(test-recurrent-state-rollback.cpp) +llama_build(test-state-seq-on-device.cpp) if (NOT WIN32 OR NOT BUILD_SHARED_LIBS) # these tests are disabled on Windows because they use internal functions not exported with LLAMA_API (when building with shared libraries) @@ -219,6 +220,15 @@ if (NOT WIN32 OR NOT BUILD_SHARED_LIBS) FIXTURES_REQUIRED generate-models ) + llama_test( + test-state-seq-on-device + LABEL main + ARGS "${MODEL_DIR}/llama-dense.gguf" + ) + set_tests_properties(test-state-seq-on-device PROPERTIES + FIXTURES_REQUIRED generate-models + ) + llama_test( test-recurrent-state-rollback NAME test-recurrent-state-rollback-nemotron-h diff --git a/tests/test-state-seq-on-device.cpp b/tests/test-state-seq-on-device.cpp new file mode 100644 index 000000000000..445ea74934a2 --- /dev/null +++ b/tests/test-state-seq-on-device.cpp @@ -0,0 +1,150 @@ +#include "llama.h" + +#include +#include +#include +#include +#include +#include + +static size_t set_on_device(llama_context * ctx, const uint8_t * src, size_t size, llama_seq_id dest_seq_id, bool * escaped) { + *escaped = false; + try { + return llama_state_seq_set_data_ext(ctx, src, size, dest_seq_id, LLAMA_STATE_SEQ_FLAGS_ON_DEVICE); + } catch (const std::exception & err) { + fprintf(stderr, "%s : exception escaped C API: %s\n", __func__, err.what()); + *escaped = true; + return 0; + } catch (...) { + fprintf(stderr, "%s : exception escaped C API\n", __func__); + *escaped = true; + return 0; + } +} + +static bool expect_fail(llama_context * ctx, const uint8_t * src, size_t size, const char * desc) { + bool escaped = false; + const size_t nset = set_on_device(ctx, src, size, 0, &escaped); + if (escaped) { + fprintf(stderr, "%s : %s\n", __func__, desc); + return false; + } + if (nset != 0) { + fprintf(stderr, "%s : expected 0 for %s, got %zu\n", __func__, desc, nset); + return false; + } + return true; +} + +int main(int argc, char ** argv) { + if (argc < 2) { + fprintf(stderr, "usage: %s \n", argv[0]); + return 1; + } + + const std::string model_path = argv[1]; + + llama_backend_init(); + + llama_model_params mparams = llama_model_default_params(); + mparams.n_gpu_layers = 0; + + llama_model * model = llama_model_load_from_file(model_path.c_str(), mparams); + if (model == nullptr) { + fprintf(stderr, "%s : failed to load the model\n", __func__); + return 1; + } + + llama_context_params cparams = llama_context_default_params(); + cparams.n_ctx = 256; + cparams.n_batch = 64; + cparams.n_seq_max = 4; + + llama_context * ctx = llama_init_from_model(model, cparams); + if (ctx == nullptr) { + fprintf(stderr, "%s : failed to create the context\n", __func__); + llama_model_free(model); + return 1; + } + + const llama_vocab * vocab = llama_model_get_vocab(model); + const int n_vocab = llama_vocab_n_tokens(vocab); + + std::vector tokens; + for (int i = 0; i < 8; ++i) { + tokens.push_back(1 + i % (n_vocab - 1)); + } + + if (llama_decode(ctx, llama_batch_get_one(tokens.data(), tokens.size())) != 0) { + fprintf(stderr, "%s : failed to decode\n", __func__); + llama_free(ctx); + llama_model_free(model); + return 1; + } + + const llama_state_seq_flags flags = LLAMA_STATE_SEQ_FLAGS_ON_DEVICE; + std::vector state(llama_state_seq_get_size_ext(ctx, 0, flags)); + if (state.size() < 8) { + fprintf(stderr, "%s : on-device state is too small: %zu\n", __func__, state.size()); + llama_free(ctx); + llama_model_free(model); + return 1; + } + + if (llama_state_seq_get_data_ext(ctx, state.data(), state.size(), 0, flags) != state.size()) { + fprintf(stderr, "%s : failed to save the state\n", __func__); + llama_free(ctx); + llama_model_free(model); + return 1; + } + + bool escaped = false; + const size_t nset_ok = set_on_device(ctx, state.data(), state.size(), 0, &escaped); + if (escaped || nset_ok != state.size()) { + fprintf(stderr, "%s : valid on-device restore failed: returned %zu, expected %zu\n", + __func__, nset_ok, state.size()); + llama_free(ctx); + llama_model_free(model); + return 1; + } + + { + std::vector buf = state; + const uint32_t magic_bad = 0xdeadbeef; + memcpy(buf.data(), &magic_bad, sizeof(magic_bad)); + if (!expect_fail(ctx, buf.data(), buf.size(), "wrong magic")) { + llama_free(ctx); + llama_model_free(model); + return 1; + } + } + + if (!expect_fail(ctx, state.data(), 2, "buffer shorter than magic")) { + llama_free(ctx); + llama_model_free(model); + return 1; + } + + if (!expect_fail(ctx, state.data(), 6, "truncated seq_id")) { + llama_free(ctx); + llama_model_free(model); + return 1; + } + + { + std::vector buf = state; + const llama_seq_id seq_id_bad = 3; + memcpy(buf.data() + sizeof(uint32_t), &seq_id_bad, sizeof(seq_id_bad)); + if (!expect_fail(ctx, buf.data(), buf.size(), "unknown source seq_id")) { + llama_free(ctx); + llama_model_free(model); + return 1; + } + } + + llama_free(ctx); + llama_model_free(model); + llama_backend_free(); + + return 0; +}