diff --git a/src/llama-memory-recurrent.cpp b/src/llama-memory-recurrent.cpp index a7f9263115..8353a30932 100644 --- a/src/llama-memory-recurrent.cpp +++ b/src/llama-memory-recurrent.cpp @@ -125,6 +125,13 @@ llama_memory_recurrent::llama_memory_recurrent( ctxs_bufs.emplace_back(std::move(ctx), buf); } + if (is_empty()) { + if (n_rs_seq > 0) { + n_rs_seq = 0; + LLAMA_LOG_INFO("%s: disabling rollback snapshots because the memory module is empty\n", __func__); + } + } + { const size_t memory_size_r = size_r_bytes(); const size_t memory_size_s = size_s_bytes(); @@ -193,7 +200,7 @@ bool llama_memory_recurrent::seq_rm(llama_seq_id seq_id, llama_pos p0, llama_pos // partial rollback via per-token snapshot index (bounded by n_rs_seq) if (0 < p0 && p0 <= cell.pos && p1 > cell.pos) { // the filter kept no layer (e.g. an MTP draft context), so only the position moves back - if (ctxs_bufs.empty()) { + if (is_empty()) { cell.pos = p0 - 1; return true; } @@ -723,6 +730,11 @@ bool llama_memory_recurrent::get_can_shift() const { return true; } +bool llama_memory_recurrent::is_empty() const { + assert(total_size() == 0); + return ctxs_bufs.empty(); +} + size_t llama_memory_recurrent::total_size() const { size_t size = 0; for (const auto & [_, buf] : ctxs_bufs) { diff --git a/src/llama-memory-recurrent.h b/src/llama-memory-recurrent.h index 25ade10e51..08489a4bee 100644 --- a/src/llama-memory-recurrent.h +++ b/src/llama-memory-recurrent.h @@ -123,6 +123,9 @@ private: // ggml contexts for the KV cache along with the allocated backend buffers: std::vector> ctxs_bufs; + // true if no layers - can happen if the layer filter removes all layers + bool is_empty() const; + size_t total_size() const; size_t size_r_bytes() const;