mirror of
https://github.com/ggml-org/llama.cpp.git
synced 2026-09-27 21:46:57 +02:00
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
This commit is contained in:
@@ -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;
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,8 +1,10 @@
|
||||
#include "llama-impl.h"
|
||||
|
||||
#include "ggml-backend.h"
|
||||
#include "gguf.h"
|
||||
#include "llama.h"
|
||||
|
||||
#include <algorithm>
|
||||
#include <cinttypes>
|
||||
#include <climits>
|
||||
#include <cstdarg>
|
||||
@@ -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<uint8_t> 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;
|
||||
|
||||
@@ -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 ?
|
||||
|
||||
@@ -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;
|
||||
|
||||
|
||||
@@ -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 {
|
||||
|
||||
+112
-5
@@ -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
|
||||
//
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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
|
||||
//
|
||||
|
||||
@@ -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 {
|
||||
|
||||
Reference in New Issue
Block a user