kv-cache : fix M-RoPE checkpoints (#20132)
This commit is contained in:
+3
-1
@@ -394,11 +394,13 @@ llama_ubatch llama_batch_allocr::ubatch_reserve(uint32_t n_seq_tokens, uint32_t
|
|||||||
clear();
|
clear();
|
||||||
split_reset();
|
split_reset();
|
||||||
|
|
||||||
|
const int64_t n_pos_all = (int64_t) n_tokens*n_pos_per_embd;
|
||||||
|
|
||||||
auto udata = std::make_shared<llama_ubatch::data_t>();
|
auto udata = std::make_shared<llama_ubatch::data_t>();
|
||||||
|
|
||||||
udata->token .resize(n_tokens);
|
udata->token .resize(n_tokens);
|
||||||
udata->embd .clear();
|
udata->embd .clear();
|
||||||
udata->pos .resize(n_tokens);
|
udata->pos .resize(n_pos_all);
|
||||||
udata->n_seq_id .resize(n_tokens);
|
udata->n_seq_id .resize(n_tokens);
|
||||||
udata->seq_id .resize(n_tokens);
|
udata->seq_id .resize(n_tokens);
|
||||||
udata->seq_id_unq.resize(0);
|
udata->seq_id_unq.resize(0);
|
||||||
|
|||||||
+12
-2
@@ -1760,8 +1760,10 @@ void llama_kv_cache::state_write_meta(llama_io_write_i & io, const cell_ranges_t
|
|||||||
io.write(&pos, sizeof(pos));
|
io.write(&pos, sizeof(pos));
|
||||||
io.write(&n_seq_id, sizeof(n_seq_id));
|
io.write(&n_seq_id, sizeof(n_seq_id));
|
||||||
|
|
||||||
// TODO: we also need to save llama_kv_cell_ext when apply_ubatch() support loading it
|
if (hparams.n_pos_per_embd() > 1) {
|
||||||
// see: https://github.com/ggml-org/llama.cpp/pull/16825#issuecomment-3460868350
|
const llama_kv_cell_ext ext = cells.ext_get(i);
|
||||||
|
io.write(&ext, sizeof(ext));
|
||||||
|
}
|
||||||
|
|
||||||
for (const auto & seq_id : seq_ids) {
|
for (const auto & seq_id : seq_ids) {
|
||||||
io.write(&seq_id, sizeof(seq_id));
|
io.write(&seq_id, sizeof(seq_id));
|
||||||
@@ -1895,6 +1897,14 @@ bool llama_kv_cache::state_read_meta(llama_io_read_i & io, uint32_t strm, uint32
|
|||||||
return false;
|
return false;
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if (hparams.n_pos_per_embd() > 1) {
|
||||||
|
llama_kv_cell_ext ext;
|
||||||
|
io.read_to(&ext, sizeof(ext));
|
||||||
|
|
||||||
|
ubatch.pos[i + ubatch.n_tokens] = ext.y;
|
||||||
|
ubatch.pos[i + ubatch.n_tokens*2] = ext.x;
|
||||||
|
}
|
||||||
|
|
||||||
// read the sequence id, but directly discard it - we will use dest_seq_id instead
|
// read the sequence id, but directly discard it - we will use dest_seq_id instead
|
||||||
{
|
{
|
||||||
llama_seq_id seq_id;
|
llama_seq_id seq_id;
|
||||||
|
|||||||
Reference in New Issue
Block a user