From 08618ff8e735141d8e4e5be28e6d6af170e4757b Mon Sep 17 00:00:00 2001 From: Chipmunk <101038159+CHIPMUNK-T0T@users.noreply.github.com> Date: Sat, 26 Sep 2026 16:23:03 +0900 Subject: [PATCH] llama : fix K/V and recurrent state cleanup after failed restores (#27530) * llama : add discard for deferred state writes * llama : add tensor zeroing helper for backends without tensor memset * llama : clear K/V data after failed sequence restore * llama : clear recurrent state data after failed sequence restore * llama : simplify discard and restore cleanup * llama : report error when abnormal cell count is found in state_read_meta * llama : clear attention state on hybrid restore failure * tests : cover failed state restore cleanup * llama : clear MLA state on dsa restore failure * tests : update test for rebased test suite * llama : clarify comment in llama_memory_recurrent::state_read --- src/llama-context.cpp | 11 +++ src/llama-impl.cpp | 12 +++ src/llama-impl.h | 2 + src/llama-io.h | 3 + src/llama-kv-cache-dsa.cpp | 10 +- src/llama-kv-cache.cpp | 117 ++++++++++++++++++++++- src/llama-kv-cache.h | 5 + src/llama-memory-hybrid.cpp | 16 +++- src/llama-memory-recurrent.cpp | 54 +++++++++-- src/llama-memory-recurrent.h | 2 + tests/test-save-load-state.cpp | 168 ++++++++++++++++++++++++++++++++- 11 files changed, 383 insertions(+), 17 deletions(-) diff --git a/src/llama-context.cpp b/src/llama-context.cpp index 8675f60873..99e55da687 100644 --- a/src/llama-context.cpp +++ b/src/llama-context.cpp @@ -2736,6 +2736,10 @@ public: buf_size -= size; } + void discard() override { + rinfos.clear(); + } + size_t n_bytes() override { return size_read; } @@ -3086,6 +3090,11 @@ public: rinfos.push_back({tensor, ptr, size, offset}); } + void discard() override { + rinfos.clear(); + buf_size = 0; + } + size_t n_bytes() override { return size_read; } @@ -3132,6 +3141,7 @@ size_t llama_context::state_set_data(const uint8_t * src, size_t size) { return state_read_data(io); } catch (const std::exception & err) { LLAMA_LOG_ERROR("%s: error loading state: %s\n", __func__, err.what()); + io.discard(); return 0; } } @@ -3205,6 +3215,7 @@ size_t llama_context::state_seq_set_data(llama_seq_id seq_id, const uint8_t * sr return state_seq_read_data(*io, seq_id, flags); } catch (const std::exception & err) { LLAMA_LOG_ERROR("%s: error loading state: %s\n", __func__, err.what()); + io->discard(); return 0; } } diff --git a/src/llama-impl.cpp b/src/llama-impl.cpp index b3a94b946d..5ec400e96d 100644 --- a/src/llama-impl.cpp +++ b/src/llama-impl.cpp @@ -1,8 +1,10 @@ #include "llama-impl.h" +#include "ggml-backend.h" #include "gguf.h" #include "llama.h" +#include #include #include #include @@ -66,6 +68,16 @@ void llama_log_callback_default(ggml_log_level level, const char * text, void * fflush(stderr); } +void llama_clear_tensor_data(ggml_tensor * t, size_t offset, size_t size) { + static const std::vector zeros(1024*1024, 0); + + // not all backend buffers implement ggml_backend_tensor_memset(), so write zeros instead + // TODO: make this a generic fallback in `ggml_backend_tensor_memset` when `set_tensor` is available + for (size_t ofs = 0; ofs < size; ofs += zeros.size()) { + ggml_backend_tensor_set(t, zeros.data(), offset + ofs, std::min(size - ofs, zeros.size())); + } +} + void replace_all(std::string & s, const std::string & search, const std::string & replace) { if (search.empty()) { return; diff --git a/src/llama-impl.h b/src/llama-impl.h index 4988b06d2c..c34a6473b7 100644 --- a/src/llama-impl.h +++ b/src/llama-impl.h @@ -93,6 +93,8 @@ struct buffer_view { } }; +void llama_clear_tensor_data(ggml_tensor * t, size_t offset, size_t size); + void replace_all(std::string & s, const std::string & search, const std::string & replace); // TODO: rename to llama_format ? diff --git a/src/llama-io.h b/src/llama-io.h index f276af4fb9..0aca733107 100644 --- a/src/llama-io.h +++ b/src/llama-io.h @@ -28,6 +28,9 @@ public: virtual void read(void * dst, size_t size) = 0; virtual void read_tensor(ggml_tensor * tensor, size_t offset, size_t size) = 0; + // drop tensor data that has been read but not yet applied (e.g. when a restore fails) + virtual void discard() {} + // bytes read so far virtual size_t n_bytes() = 0; diff --git a/src/llama-kv-cache-dsa.cpp b/src/llama-kv-cache-dsa.cpp index 96cb045d2e..6e5f7d4840 100644 --- a/src/llama-kv-cache-dsa.cpp +++ b/src/llama-kv-cache-dsa.cpp @@ -168,7 +168,15 @@ void llama_kv_cache_dsa::state_write(llama_io_write_i & io, llama_seq_id seq_id, void llama_kv_cache_dsa::state_read(llama_io_read_i & io, llama_seq_id seq_id, llama_state_seq_flags flags) { kv_mla->state_read(io, seq_id, flags); - kv_lid->state_read(io, seq_id, flags); + + try { + kv_lid->state_read(io, seq_id, flags); + } catch (...) { + // the MLA part is already restored - undo it, so that a failed restore leaves nothing behind + kv_mla->state_clear(seq_id); + + throw; + } } llama_kv_cache * llama_kv_cache_dsa::get_mla() const { diff --git a/src/llama-kv-cache.cpp b/src/llama-kv-cache.cpp index 332d1abe02..1c93f908b0 100644 --- a/src/llama-kv-cache.cpp +++ b/src/llama-kv-cache.cpp @@ -2188,11 +2188,7 @@ const slot_info_vec_t * sinfos_in) { } if (!res) { - if (seq_id == -1) { - clear(true); - } else { - seq_rm(seq_id, -1, -1); - } + state_clear(seq_id, strm, sinfo); throw std::runtime_error("failed to restore kv cache"); } @@ -2340,6 +2336,11 @@ bool llama_kv_cache::state_read_meta(llama_io_read_i & io, uint32_t strm, uint32 if (dest_seq_id != -1) { // single sequence + if (cell_count > cells.size()) { + LLAMA_LOG_ERROR("%s: not enough cells in kv cache\n", __func__); + return false; + } + seq_rm(dest_seq_id, -1, -1); llama_batch_allocr balloc(hparams.n_pos_per_embd()); @@ -2663,6 +2664,112 @@ bool llama_kv_cache::state_read_data(llama_io_read_i & io, uint32_t strm, uint32 return true; } +void llama_kv_cache::state_clear(llama_seq_id seq_id) { + if (seq_id == -1) { + clear(true); + return; + } + + GGML_ASSERT(seq_id >= 0 && (size_t) seq_id < seq_to_stream.size()); + + const uint32_t strm = seq_to_stream[seq_id]; + + const auto & cells = v_cells[strm]; + + slot_info sinfo; + sinfo.s0 = strm; + sinfo.s1 = strm; + sinfo.resize(1); + sinfo.strm[0] = strm; + + // a cell that another sequence still uses keeps its data + for (uint32_t i = 0; i < cells.size(); ++i) { + if (cells.seq_has(i, seq_id) && cells.seq_count(i) == 1) { + sinfo.idxs[0].push_back(i); + } + } + + state_clear(seq_id, strm, sinfo); +} + +// the cleared ranges mirror the write pattern of state_read_data() - keep both in sync +void llama_kv_cache::state_clear(llama_seq_id seq_id, uint32_t strm, const slot_info & sinfo) { + if (seq_id == -1) { + clear(true); + return; + } + + seq_rm(seq_id, -1, -1); + + // zero the K/V data of the failed restore attempt - the attention can still read the data of free cells + if (sinfo.empty() || sinfo.size() == 0) { + return; + } + + const auto & cells = v_cells[strm]; + + const uint32_t cell_count = sinfo.size(); + + const bool is_contiguous = sinfo.is_contiguous(); + + for (const auto & layer : layers) { + const uint32_t il = layer.il; + + const uint32_t n_embd_k_gqa = hparams.n_embd_k_gqa(il); + + auto * k = layer.k_stream[strm]; + + const size_t k_size_row = ggml_row_size(k->type, n_embd_k_gqa); + + if (is_contiguous) { + llama_clear_tensor_data(k, sinfo.head() * k_size_row, cell_count * k_size_row); + } else { + for (uint32_t i = 0; i < cell_count; ++i) { + llama_clear_tensor_data(k, sinfo.idxs[0][i] * k_size_row, k_size_row); + } + } + } + + for (const auto & layer : layers) { + const uint32_t il = layer.il; + + const uint32_t n_embd_v_gqa = hparams.n_embd_v_gqa(il); + + auto * v = layer.v_stream[strm]; + if (!v) { + continue; + } + + if (!v_trans) { + const size_t v_size_row = ggml_row_size(v->type, n_embd_v_gqa); + + if (is_contiguous) { + llama_clear_tensor_data(v, sinfo.head() * v_size_row, cell_count * v_size_row); + } else { + for (uint32_t i = 0; i < cell_count; ++i) { + llama_clear_tensor_data(v, sinfo.idxs[0][i] * v_size_row, v_size_row); + } + } + } else { + const size_t v_size_el = ggml_type_size(v->type); + + if (is_contiguous) { + const uint32_t h = sinfo.head(); + + for (uint32_t j = 0; j < n_embd_v_gqa; ++j) { + llama_clear_tensor_data(v, (h + j * cells.size()) * v_size_el, cell_count * v_size_el); + } + } else { + for (uint32_t j = 0; j < n_embd_v_gqa; ++j) { + for (uint32_t i = 0; i < cell_count; ++i) { + llama_clear_tensor_data(v, (sinfo.idxs[0][i] + j * cells.size()) * v_size_el, v_size_el); + } + } + } + } + } +} + // // llama_kv_cache_context // diff --git a/src/llama-kv-cache.h b/src/llama-kv-cache.h index c4d8699def..5051f43433 100644 --- a/src/llama-kv-cache.h +++ b/src/llama-kv-cache.h @@ -179,6 +179,9 @@ public: slot_info_vec_t * sinfos_out, const slot_info_vec_t * sinfos_in); + // undo a state_read() of seq_id (-1 for the whole cache) that another memory module failed to complete + void state_clear(llama_seq_id seq_id); + // // graph_build API // @@ -345,6 +348,8 @@ private: // sinfo_in, when set, replaces the find_slot call: the cells are given by the caller bool state_read_meta(llama_io_read_i & io, uint32_t strm, uint32_t cell_count, slot_info & sinfo, llama_seq_id dest_seq_id = -1, const slot_info * sinfo_in = nullptr); bool state_read_data(llama_io_read_i & io, uint32_t strm, uint32_t cell_count, const slot_info & sinfo); + + void state_clear(llama_seq_id seq_id, uint32_t strm, const slot_info & sinfo); }; class llama_kv_cache_context : public llama_memory_context_i { diff --git a/src/llama-memory-hybrid.cpp b/src/llama-memory-hybrid.cpp index 42c7381a9e..62f12de954 100644 --- a/src/llama-memory-hybrid.cpp +++ b/src/llama-memory-hybrid.cpp @@ -195,10 +195,22 @@ void llama_memory_hybrid::state_write(llama_io_write_i & io, llama_seq_id seq_id } void llama_memory_hybrid::state_read(llama_io_read_i & io, llama_seq_id seq_id, llama_state_seq_flags flags) { - if ((flags & LLAMA_STATE_SEQ_FLAGS_PARTIAL_ONLY) == 0) { + const bool read_attn = (flags & LLAMA_STATE_SEQ_FLAGS_PARTIAL_ONLY) == 0; + + if (read_attn) { mem_attn->state_read(io, seq_id, flags); } - mem_recr->state_read(io, seq_id, flags); + + try { + mem_recr->state_read(io, seq_id, flags); + } catch (...) { + // the attention part is already restored - undo it + if (read_attn) { + mem_attn->state_clear(seq_id); + } + + throw; + } } llama_kv_cache * llama_memory_hybrid::get_mem_attn() const { diff --git a/src/llama-memory-recurrent.cpp b/src/llama-memory-recurrent.cpp index 57919accf0..528c90c419 100644 --- a/src/llama-memory-recurrent.cpp +++ b/src/llama-memory-recurrent.cpp @@ -852,7 +852,12 @@ void llama_memory_recurrent::state_read(llama_io_read_i & io, llama_seq_id seq_i bool res = true; - res = res && state_read_meta(io, cell_count, seq_id); + // save the head of the restored cells - could be needed to clear the state + // the head is valid only when state_read_meta() succeeded + const bool meta_read = state_read_meta(io, cell_count, seq_id); + const uint32_t cell_head = head; + + res = res && meta_read; try { res = res && state_read_data(io, cell_count); @@ -861,12 +866,7 @@ void llama_memory_recurrent::state_read(llama_io_read_i & io, llama_seq_id seq_i } if (!res) { - // TODO: fix incosistent handling of `seq_id < 0` and `seq_id == -1` in the codebase [TAG_LLAMA_SEQ_ID_NEG] - if (seq_id == -1) { - clear(true); - } else { - seq_rm(seq_id, -1, -1); - } + state_clear(seq_id, cell_head, meta_read ? cell_count : 0); throw std::runtime_error("failed to restore kv cache"); } @@ -992,6 +992,11 @@ void llama_memory_recurrent::state_write_data(llama_io_write_i & io, const std:: bool llama_memory_recurrent::state_read_meta(llama_io_read_i & io, uint32_t cell_count, llama_seq_id dest_seq_id) { if (dest_seq_id != -1) { // single sequence + if (cell_count > size) { + LLAMA_LOG_ERROR("%s: not enough cells in kv cache\n", __func__); + return false; + } + seq_rm(dest_seq_id, -1, -1); if (cell_count == 0) { @@ -1223,6 +1228,41 @@ bool llama_memory_recurrent::state_read_data(llama_io_read_i & io, uint32_t cell return true; } +// the cleared ranges mirror the write pattern of state_read_data() - keep both in sync +// the transposed s layout is not handled - state_read_data() rejects it before any write +void llama_memory_recurrent::state_clear(llama_seq_id seq_id, uint32_t cell_head, uint32_t cell_count) { + // TODO: fix incosistent handling of `seq_id < 0` and `seq_id == -1` in the codebase [TAG_LLAMA_SEQ_ID_NEG] + if (seq_id == -1) { + clear(true); + return; + } + + seq_rm(seq_id, -1, -1); + + if (cell_count == 0) { + return; + } + + const uint32_t n_layer = hparams.n_layer(); + + for (uint32_t il = 0; il < n_layer; ++il) { + if (r_l[il] != nullptr) { + const size_t r_size_row = ggml_row_size(r_l[il]->type, hparams.n_embd_r()); + llama_clear_tensor_data(r_l[il], cell_head * r_size_row, cell_count * r_size_row); + } + + if (s_l[il] != nullptr) { + const size_t s_size_row = ggml_row_size(s_l[il]->type, hparams.n_embd_s()); + llama_clear_tensor_data(s_l[il], cell_head * s_size_row, cell_count * s_size_row); + } + + if (p_l[il] != nullptr) { + const size_t p_size_row = ggml_row_size(p_l[il]->type, hparams.ple_conv_state()); + llama_clear_tensor_data(p_l[il], cell_head * p_size_row, cell_count * p_size_row); + } + } +} + // // llama_memory_recurrent_context // diff --git a/src/llama-memory-recurrent.h b/src/llama-memory-recurrent.h index 4abb3f5cf5..25ade10e51 100644 --- a/src/llama-memory-recurrent.h +++ b/src/llama-memory-recurrent.h @@ -134,6 +134,8 @@ private: bool state_read_meta(llama_io_read_i & io, uint32_t cell_count, llama_seq_id dest_seq_id = -1); bool state_read_data(llama_io_read_i & io, uint32_t cell_count); + + void state_clear(llama_seq_id seq_id, uint32_t cell_head, uint32_t cell_count); }; class llama_memory_recurrent_context : public llama_memory_context_i { diff --git a/tests/test-save-load-state.cpp b/tests/test-save-load-state.cpp index 33ee7143af..083ad6a6ca 100644 --- a/tests/test-save-load-state.cpp +++ b/tests/test-save-load-state.cpp @@ -5,10 +5,16 @@ #include #include +#include +#include #include #include +#include +#include +#include #include #include +#include #include constexpr double NMSE_THRESHOLD = 1e-5; @@ -599,6 +605,161 @@ static bool test_state_roundtrip(struct llama_model * model, const struct common } +// overwrite the tensor data with 0xff bytes (NaN when read as f16/f32), so that the restore fails +static bool corrupt_state(std::vector & data) { + if (data.size() < 3*4096) { + LOG_ERR("%s: state of %zu bytes is too small to corrupt\n", __func__, data.size()); + return false; + } + + std::fill(data.begin() + 4096, data.end() - data.size()/4, 0xff); + return true; +} + + +// Test 9: state restore failure +// a failed restore must leave the sequence empty and must not change the logits of other sequences +static bool test_state_restore_failure(struct llama_model * model, const struct common_params & params, const llama_tokens & tokens) { + auto params_ctx = common_context_params_to_llama(params); + params_ctx.n_ctx = 256; + params_ctx.n_seq_max = 4; + params_ctx.kv_unified = true; + + // without flash attention, corrupted data left behind by the restore shows up as NaN logits on the other sequences + params_ctx.flash_attn_type = LLAMA_FLASH_ATTN_TYPE_DISABLED; + + auto ctx = llama_context_ptr{llama_init_from_model(model, params_ctx)}; + if (!ctx) { + LOG_ERR("%s: failed to create context\n", __func__); + return false; + } + + LOGV(LOG_LEVEL_INFO, "\n=== Test 9: state restore failure ===\n"); + + llama_memory_t mem = llama_get_memory(ctx.get()); + if (mem == nullptr) { + LOGV(LOG_LEVEL_INFO, "PASS (model has no memory)\n"); + return true; + } + + const auto decode = [&](const llama_tokens & inp, llama_seq_id seq_id, std::vector * logits_out) { + llama_batch_ptr batch(inp.size(), 0, 1); + for (size_t i = 0; i < inp.size(); ++i) { + common_batch_add(batch.get(), inp[i], i, { seq_id }, i == inp.size() - 1); + } + + if (llama_decode(ctx.get(), batch.get())) { + LOG_ERR("%s: failed to decode on sequence %d\n", __func__, seq_id); + return false; + } + + if (logits_out && !get_current_logits(ctx.get(), *logits_out)) { + LOG_ERR("%s: failed to get logits\n", __func__); + return false; + } + + return true; + }; + + const llama_tokens tokens_save (tokens.begin(), tokens.begin() + std::min(24, tokens.size())); + const llama_tokens tokens_verify(tokens.end() - std::min(8, tokens.size()), tokens.end()); + + // the registered tests share a working directory, so the state file is named after the model + const std::string path = "state-restore-failure." + std::filesystem::path(params.model.path).filename().string() + ".tmp.bin"; + + llama_memory_clear(mem, true); + + std::vector baseline; + if (!decode(tokens_verify, 1, &baseline)) { + return false; + } + + const std::vector>> cases = { + { "buffer", [&]() { + std::vector state(llama_state_seq_get_size(ctx.get(), 0)); + GGML_ASSERT(llama_state_seq_get_data(ctx.get(), state.data(), state.size(), 0) == state.size()); + llama_memory_seq_rm(mem, 0, -1, -1); + + if (!corrupt_state(state)) { + return false; + } + + return llama_state_seq_set_data(ctx.get(), state.data(), state.size(), 0) == 0; + }}, + { "file", [&]() { + GGML_ASSERT(llama_state_seq_save_file(ctx.get(), path.c_str(), 0, tokens_save.data(), tokens_save.size()) > 0); + llama_memory_seq_rm(mem, 0, -1, -1); + + std::vector data; + { + std::ifstream f(path, std::ios::binary); + data.assign(std::istreambuf_iterator(f), std::istreambuf_iterator()); + } + + if (!corrupt_state(data)) { + std::remove(path.c_str()); + return false; + } + + { + std::ofstream f(path, std::ios::binary); + f.write((const char *) data.data(), data.size()); + } + + llama_tokens tokens_out(tokens_save.size()); + size_t n_token_count = 0; + const size_t nread = llama_state_seq_load_file(ctx.get(), path.c_str(), 0, tokens_out.data(), tokens_out.size(), &n_token_count); + std::remove(path.c_str()); + + return nread == 0; + }}, + }; + + for (const auto & [name, restore_failed] : cases) { + llama_memory_clear(mem, true); + + if (!decode(tokens_save, 0, nullptr)) { + return false; + } + + if (!restore_failed()) { + LOG_ERR("%s: %s: restoring a corrupted state did not fail\n", __func__, name); + return false; + } + + if (llama_memory_seq_pos_max(mem, 0) != -1) { + LOG_ERR("%s: %s: sequence not empty after failed restore\n", __func__, name); + return false; + } + + std::vector logits; + if (!decode(tokens_verify, 1, &logits)) { + return false; + } + + float diff_max = 0.0f; + size_t n_nan = 0; + for (size_t i = 0; i < logits.size(); ++i) { + if (std::isnan(logits[i]) || std::isnan(baseline[i])) { + n_nan++; + } else { + diff_max = std::max(diff_max, std::fabs(logits[i] - baseline[i])); + } + } + + if (n_nan > 0 || diff_max > 1e-6f) { + LOG_ERR("%s: %s: logits changed after failed restore (max diff = %g, nan = %zu)\n", __func__, name, diff_max, n_nan); + return false; + } + + LOG_TRC("%s: %s: logits match (max diff = %g)\n", __func__, name, diff_max); + } + + LOGV(LOG_LEVEL_INFO, "\nPASS\n"); + return true; +} + + struct test_suite { std::vector results; @@ -609,10 +770,10 @@ struct test_suite { // column headers for the --models table, one per test, in the order they are run static const std::vector test_names = { - "baseline", "seq_rm", "state_load", "cp_h", "cp_d", "cp_h_s", "cp_d_s", "rt", + "baseline", "seq_rm", "state_load", "cp_h", "cp_d", "cp_h_s", "cp_d_s", "rt", "rf", }; -// Run the full save/load test suite (tests 1-8) for a single model. +// Run the full save/load test suite (tests 1-9) for a single model. // Returns the per-test results. static test_suite run_save_load_tests_for_model(const std::string & model_path, const struct common_params & base_params) { test_suite suite; @@ -688,6 +849,9 @@ static test_suite run_save_load_tests_for_model(const std::string & model_path, // Test 8: state blob round-trip suite.results.push_back(test_state_roundtrip(model, params, tokens) ? test_status::PASS : test_status::FAIL); + // Test 9: state restore failure + suite.results.push_back(test_state_restore_failure(model, params, tokens) ? test_status::PASS : test_status::FAIL); + return suite; }