Skip to content

Commit ca0c575

Browse files
Stratumclaude
andcommitted
feat(server): reuse restored KV/recurrent state instead of reprocessing
On hybrid recurrent models, slot restore followed by prompt reuse always forced a full reprocess: only the current recurrent plane was persisted, rs_idx was reset on read, and the reuse gate reset n_past whenever the aggregate pos_min came from recurrent memory. This made cross-process KV restore useless for skipping prompt processing. - llama-memory-recurrent: persist all 1+n_rs_seq rollback planes (offset k*size) and rs_idx in state_write_data/state_read_data, so the running state can roll back to the reuse boundary after a restore. - server-context: in the reuse gate, when the memory supports bounded recurrent rollback and (pos_min - n_past + 1) <= n_rs_seq, keep n_past and reuse; otherwise fall back to full reprocess (no abort). Add a .dft sidecar to save/restore the speculative draft context so seq_rm on the draft does not abort under --spec-type. - common: add KVARN_N_RS_SEQ env override for the recurrent rollback budget (native comes from speculative need_n_rs_seq(); 0 without speculation). Result: restored state is reused instead of reprocessed (prompt_n 4843->19 in tests), correct recall, works cross-GPU, with safe reprocess fallback. Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com> Claude-Session: https://claude.ai/code/session_01RErC5TKNN678KpqAQwFatc
1 parent f16b973 commit ca0c575

3 files changed

Lines changed: 75 additions & 18 deletions

File tree

‎common/common.cpp‎

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1697,6 +1697,7 @@ struct llama_context_params common_context_params_to_llama(const common_params &
16971697
cparams.n_ctx = params.n_ctx;
16981698
cparams.n_seq_max = params.n_parallel;
16991699
cparams.n_rs_seq = params.speculative.need_n_rs_seq();
1700+
if (const char * kvenv = getenv("KVARN_N_RS_SEQ")) { cparams.n_rs_seq = (uint32_t) atoi(kvenv); } // KVSNAP test override
17001701
cparams.n_outputs_max = std::max(params.n_outputs_max, 0);
17011702
cparams.n_outputs_max_per_seq = std::max(params.n_outputs_max_per_seq, 0);
17021703
cparams.n_batch = params.n_batch;

‎src/llama-memory-recurrent.cpp‎

Lines changed: 32 additions & 14 deletions
Original file line numberDiff line numberDiff line change
@@ -807,7 +807,16 @@ void llama_memory_recurrent::state_write(llama_io_write_i & io, llama_seq_id seq
807807
io.write(&cell_count, sizeof(cell_count));
808808

809809
state_write_meta(io, cell_ranges, seq_id);
810-
state_write_data(io, cell_ranges_data);
810+
811+
// KVSNAP: persist the current rollback-plane index so read can restore snapshots
812+
uint32_t rs_idx_save = 0;
813+
if (n_rs_seq != 0 && seq_id >= 0 && (size_t) seq_id < rs_idx.size()) {
814+
rs_idx_save = rs_idx[seq_id];
815+
}
816+
io.write(&rs_idx_save, sizeof(rs_idx_save));
817+
818+
(void) cell_ranges_data;
819+
state_write_data(io, cell_ranges); // KVSNAP: base ranges; data loops all planes
811820
}
812821

813822
void llama_memory_recurrent::state_read(llama_io_read_i & io, llama_seq_id seq_id, llama_state_seq_flags flags) {
@@ -820,6 +829,9 @@ void llama_memory_recurrent::state_read(llama_io_read_i & io, llama_seq_id seq_i
820829

821830
res = res && state_read_meta(io, cell_count, seq_id);
822831

832+
uint32_t rs_idx_load = 0; // KVSNAP
833+
io.read(&rs_idx_load, sizeof(rs_idx_load));
834+
823835
try {
824836
res = res && state_read_data(io, cell_count);
825837
} catch (...) {
@@ -839,7 +851,7 @@ void llama_memory_recurrent::state_read(llama_io_read_i & io, llama_seq_id seq_i
839851
if (seq_id == -1) {
840852
std::fill(rs_idx.begin(), rs_idx.end(), 0);
841853
} else {
842-
set_rs_idx(seq_id, 0);
854+
set_rs_idx(seq_id, rs_idx_load); // KVSNAP: restore saved rollback index
843855
}
844856
}
845857
}
@@ -886,10 +898,12 @@ void llama_memory_recurrent::state_write_data(llama_io_write_i & io, const std::
886898

887899
// Write each logical cell row range. With pending recurrent rollback,
888900
// the logical current state may live in a rollback snapshot plane.
889-
for (const auto & range : cell_ranges) {
890-
const size_t range_size = range.second - range.first;
891-
const size_t buf_size = range_size * r_size_row;
892-
io.write_tensor(r_l[il], range.first * r_size_row, buf_size);
901+
for (uint32_t kp = 0; kp <= n_rs_seq; ++kp) { // KVSNAP: write all rollback planes
902+
for (const auto & range : cell_ranges) {
903+
const size_t range_size = range.second - range.first;
904+
const size_t buf_size = range_size * r_size_row;
905+
io.write_tensor(r_l[il], (range.first + kp * size) * r_size_row, buf_size);
906+
}
893907
}
894908
}
895909

@@ -908,10 +922,12 @@ void llama_memory_recurrent::state_write_data(llama_io_write_i & io, const std::
908922

909923
// Write each logical cell row range. With pending recurrent rollback,
910924
// the logical current state may live in a rollback snapshot plane.
911-
for (const auto & range : cell_ranges) {
912-
const size_t range_size = range.second - range.first;
913-
const size_t buf_size = range_size * s_size_row;
914-
io.write_tensor(s_l[il], range.first * s_size_row, buf_size);
925+
for (uint32_t kp = 0; kp <= n_rs_seq; ++kp) { // KVSNAP: write all rollback planes
926+
for (const auto & range : cell_ranges) {
927+
const size_t range_size = range.second - range.first;
928+
const size_t buf_size = range_size * s_size_row;
929+
io.write_tensor(s_l[il], (range.first + kp * size) * s_size_row, buf_size);
930+
}
915931
}
916932
}
917933
} else {
@@ -1086,8 +1102,9 @@ bool llama_memory_recurrent::state_read_data(llama_io_read_i & io, uint32_t cell
10861102
}
10871103

10881104
if (cell_count) {
1089-
// Read and set the keys for the whole cell range
1090-
io.read_tensor(r_l[il], head * r_size_row, cell_count * r_size_row);
1105+
for (uint32_t kp = 0; kp <= n_rs_seq; ++kp) { // KVSNAP: read all rollback planes
1106+
io.read_tensor(r_l[il], (head + kp * size) * r_size_row, cell_count * r_size_row);
1107+
}
10911108
}
10921109
}
10931110

@@ -1116,8 +1133,9 @@ bool llama_memory_recurrent::state_read_data(llama_io_read_i & io, uint32_t cell
11161133
}
11171134

11181135
if (cell_count) {
1119-
// Read and set the values for the whole cell range
1120-
io.read_tensor(s_l[il], head * s_size_row, cell_count * s_size_row);
1136+
for (uint32_t kp = 0; kp <= n_rs_seq; ++kp) { // KVSNAP: read all rollback planes
1137+
io.read_tensor(s_l[il], (head + kp * size) * s_size_row, cell_count * s_size_row);
1138+
}
11211139
}
11221140
}
11231141
} else {

‎tools/server/server-context.cpp‎

Lines changed: 42 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -2476,6 +2476,17 @@ struct server_context_impl {
24762476
break;
24772477
}
24782478

2479+
// KVFIX-DFT: persist the draft (speculative) context state so a restored
2480+
// slot has a consistent draft KV; else seq_rm on the draft aborts.
2481+
if (ctx_dft) {
2482+
const size_t szd = llama_state_seq_get_size_ext(ctx_dft, slot->id, LLAMA_STATE_SEQ_FLAGS_NONE);
2483+
std::vector<uint8_t> bufd(szd);
2484+
llama_state_seq_get_data_ext(ctx_dft, bufd.data(), szd, slot->id, LLAMA_STATE_SEQ_FLAGS_NONE);
2485+
const std::string dftpath = filepath + ".dft";
2486+
FILE * fdft = fopen(dftpath.c_str(), "wb");
2487+
if (fdft) { fwrite(bufd.data(), 1, szd, fdft); fclose(fdft); }
2488+
}
2489+
24792490
const int64_t t_end = ggml_time_us();
24802491
const double t_save_ms = (t_end - t_start) / 1000.0;
24812492

@@ -2535,6 +2546,21 @@ struct server_context_impl {
25352546

25362547
slot->prompt.clear();
25372548
slot->prompt.tokens = std::move(restored);
2549+
2550+
// KVFIX-DFT: restore the draft (speculative) context state to match tgt.
2551+
if (ctx_dft) {
2552+
const std::string dftpath = filepath + ".dft";
2553+
FILE * fdft = fopen(dftpath.c_str(), "rb");
2554+
if (fdft) {
2555+
fseek(fdft, 0, SEEK_END); long szd = ftell(fdft); fseek(fdft, 0, SEEK_SET);
2556+
std::vector<uint8_t> bufd(szd > 0 ? szd : 0);
2557+
const size_t rd = szd > 0 ? fread(bufd.data(), 1, (size_t) szd, fdft) : 0;
2558+
fclose(fdft);
2559+
if (rd == (size_t) szd && szd > 0) {
2560+
llama_state_seq_set_data_ext(ctx_dft, bufd.data(), bufd.size(), slot->id, LLAMA_STATE_SEQ_FLAGS_NONE);
2561+
}
2562+
}
2563+
}
25382564
} catch (const std::exception & err) {
25392565
slot->prompt_clear();
25402566
send_error(task, std::string("Unable to restore slot: ") + err.what(), ERROR_TYPE_INVALID_REQUEST);
@@ -3270,10 +3296,22 @@ struct server_context_impl {
32703296
}
32713297

32723298
if (do_reset) {
3273-
SLT_TRC(slot, "forcing full prompt re-processing due to lack of cache data (likely due to SWA or hybrid/recurrent memory, see %s)\n",
3274-
"https://github.com/ggml-org/llama.cpp/pull/13194#issuecomment-2868343055");
3275-
pos_next = 0;
3276-
n_past = 0;
3299+
// Hybrid/recurrent memory keeps only a running state, so the
3300+
// pos_min gate above wants a checkpoint. But if the memory
3301+
// supports bounded recurrent-state rollback (n_rs_seq) and the
3302+
// reuse boundary is within that budget, we can safely reuse the
3303+
// restored prefix: seq_rm rolls the running state back using the
3304+
// persisted snapshots. Otherwise fall back to full reprocess.
3305+
const int32_t n_rs = (int32_t) llama_n_rs_seq(ctx_tgt);
3306+
const llama_pos rollback = pos_min - n_past + 1;
3307+
if (n_rs > 0 && rollback <= (llama_pos) n_rs) {
3308+
SLT_DBG(slot, "recurrent reuse: rollback = %d <= n_rs_seq = %d, keeping n_past = %d\n", (int) rollback, n_rs, n_past);
3309+
} else {
3310+
SLT_TRC(slot, "forcing full prompt re-processing due to lack of cache data (likely due to SWA or hybrid/recurrent memory, see %s)\n",
3311+
"https://github.com/ggml-org/llama.cpp/pull/13194#issuecomment-2868343055");
3312+
pos_next = 0;
3313+
n_past = 0;
3314+
}
32773315
}
32783316
}
32793317
}

0 commit comments

Comments
 (0)