llama + spec: MTP Support (#22673)
* spec: support MTP * fix batch size * rename files * cont : simplify (#7) * MTP: clean-up (#9) * MTP: clean-up * review: use llama_context_type instead of llama_graph_type * review: remove llama_model_has_mtp * review: fix convert issues * convert: fix pycheck * review: formatting * use `mtp-` for identifying mtp models * convert: fix mtp conversion * mtp -> draft-mtp * remove unused llama_arch * add need_embd in speculative * llama: allow partial seq_rm for GDN models for speculative decoding Currently speculative checkpoint needs to restart from a checkpoint after some draft tokens are not accepted, this leads to some wastage in running the target again. This PR adds the ability to rollback upto `draft_max` by storing the GDN intermediates. * fix pending state * vulkan: add GDN partial rollback * meta: extend check to axis 1 * metal: add GDN partial rollback Extend the gated delta net kernel to store intermediate states for partial rollback support on the Metal backend. - Add K (snapshot slot count) as a function constant - Read input state from slot 0 of the 3D state tensor - Write intermediate states to different slots during token loop - For K=1, maintain backward-compatible single-slot behavior Ref: https://github.com/ggml-org/llama.cpp/commit/8c05923630110223669f069af2000e9cf10c02bc Assisted-by: llama.cpp:local pi * delta_net_base: use ggml_pad instead of new_tensor * review: add need_rs_seq * review: rename part_bounded to n_rs * review: deslop comments * review: rename, add asserts * server : adjust checkpoint logic (#11) * server : adjust checkpoint logic * cont : rm asserts * server-context: fix early exit * spec : fix compatibility with n-gram and add TODOs (#13) * metal : cleanup * llama : fix faulty bitwise check in recurrent memory * server : disable RS-based MTP in combination with other spec types * spec : add TODOs * cont : fix comment * cont : update comment * common : fix logic for ngram + mtp compat * llama-memory: enable checkpointing with partial rollback * cont: add test-case for loading into a dirty ctx * llama-memory-recurrent: clear rs_idx in clear * download: fix mtp path * llama-arch: fix enorm op * docs: update docs * conversion: fix type annotations --------- Co-authored-by: Georgi Gerganov <ggerganov@gmail.com>
This commit is contained in:
co-authored by
Georgi Gerganov
parent
b81c2cdd74
commit
255582687b
+104
-15
@@ -24,6 +24,7 @@ llama_memory_recurrent::llama_memory_recurrent(
|
||||
bool offload,
|
||||
uint32_t mem_size,
|
||||
uint32_t n_seq_max,
|
||||
uint32_t n_rs_seq,
|
||||
const layer_filter_cb & filter) : hparams(model.hparams), n_seq_max(n_seq_max) {
|
||||
const int32_t n_layer = hparams.n_layer;
|
||||
|
||||
@@ -31,6 +32,9 @@ llama_memory_recurrent::llama_memory_recurrent(
|
||||
size = mem_size;
|
||||
used = 0;
|
||||
|
||||
this->n_rs_seq = n_rs_seq;
|
||||
rs_idx.assign(n_seq_max, 0);
|
||||
|
||||
cells.clear();
|
||||
cells.resize(mem_size);
|
||||
|
||||
@@ -92,8 +96,9 @@ llama_memory_recurrent::llama_memory_recurrent(
|
||||
throw std::runtime_error("failed to create ggml context for rs cache");
|
||||
}
|
||||
|
||||
ggml_tensor * r = ggml_new_tensor_2d(ctx, type_r, hparams.n_embd_r(), mem_size);
|
||||
ggml_tensor * s = ggml_new_tensor_2d(ctx, type_s, hparams.n_embd_s(), mem_size);
|
||||
const uint32_t n_rows = mem_size * (1 + n_rs_seq);
|
||||
ggml_tensor * r = ggml_new_tensor_2d(ctx, type_r, hparams.n_embd_r(), n_rows);
|
||||
ggml_tensor * s = ggml_new_tensor_2d(ctx, type_s, hparams.n_embd_s(), n_rows);
|
||||
ggml_format_name(r, "cache_r_l%d", i);
|
||||
ggml_format_name(s, "cache_s_l%d", i);
|
||||
r_l[i] = r;
|
||||
@@ -115,8 +120,8 @@ llama_memory_recurrent::llama_memory_recurrent(
|
||||
const size_t memory_size_r = size_r_bytes();
|
||||
const size_t memory_size_s = size_s_bytes();
|
||||
|
||||
LLAMA_LOG_INFO("%s: size = %7.2f MiB (%6u cells, %3d layers, %2u seqs), R (%s): %7.2f MiB, S (%s): %7.2f MiB\n", __func__,
|
||||
(float)(memory_size_r + memory_size_s) / (1024.0f * 1024.0f), mem_size, n_layer, n_seq_max,
|
||||
LLAMA_LOG_INFO("%s: size = %7.2f MiB (%6u cells, %3d layers, %2u seqs %2u rs_seq), R (%s): %7.2f MiB, S (%s): %7.2f MiB\n", __func__,
|
||||
(float)(memory_size_r + memory_size_s) / (1024.0f * 1024.0f), mem_size, n_layer, n_seq_max, n_rs_seq,
|
||||
ggml_type_name(type_r), (float)memory_size_r / (1024.0f * 1024.0f),
|
||||
ggml_type_name(type_s), (float)memory_size_s / (1024.0f * 1024.0f));
|
||||
}
|
||||
@@ -138,10 +143,11 @@ void llama_memory_recurrent::clear(bool data) {
|
||||
ggml_backend_buffer_clear(buf.get(), 0);
|
||||
}
|
||||
}
|
||||
|
||||
std::fill(rs_idx.begin(), rs_idx.end(), 0);
|
||||
}
|
||||
|
||||
bool llama_memory_recurrent::seq_rm(llama_seq_id seq_id, llama_pos p0, llama_pos p1) {
|
||||
//printf("[DEBUG] calling llama_memory_recurrent::seq_rm` with `seq_id=%d, p0=%d, p1=%d`\n", seq_id, p0, p1);
|
||||
uint32_t new_head = size;
|
||||
|
||||
if (p0 < 0) {
|
||||
@@ -152,6 +158,15 @@ bool llama_memory_recurrent::seq_rm(llama_seq_id seq_id, llama_pos p0, llama_pos
|
||||
p1 = std::numeric_limits<llama_pos>::max();
|
||||
}
|
||||
|
||||
const bool rm_all = p0 == 0 && p1 == std::numeric_limits<llama_pos>::max();
|
||||
if (rm_all) {
|
||||
if (seq_id >= 0) {
|
||||
set_rs_idx(seq_id, 0);
|
||||
} else {
|
||||
std::fill(rs_idx.begin(), rs_idx.end(), 0);
|
||||
}
|
||||
}
|
||||
|
||||
// models like Mamba or RWKV can't have a state partially erased at the end
|
||||
// of the sequence because their state isn't preserved for previous tokens
|
||||
if (seq_id >= (int64_t) size) {
|
||||
@@ -161,10 +176,16 @@ bool llama_memory_recurrent::seq_rm(llama_seq_id seq_id, llama_pos p0, llama_pos
|
||||
if (0 <= seq_id) {
|
||||
int32_t & tail_id = cells[seq_id].tail;
|
||||
if (tail_id >= 0) {
|
||||
const auto & cell = cells[tail_id];
|
||||
// partial intersection is invalid if it includes the final pos
|
||||
auto & cell = cells[tail_id];
|
||||
|
||||
// partial rollback via per-token snapshot index (bounded by n_rs_seq)
|
||||
if (0 < p0 && p0 <= cell.pos && p1 > cell.pos) {
|
||||
//printf("[DEBUG] inside `llama_memory_recurrent::seq_rm`: partial intersection is invalid, so returning false, p0 = %d, cell.pos = %d, p1 = %d\n", p0, cell.pos, p1);
|
||||
const llama_pos rollback = cell.pos - (p0 - 1);
|
||||
if (rollback >= 1 && rollback <= (llama_pos) n_rs_seq) {
|
||||
set_rs_idx(seq_id, (uint32_t) rollback);
|
||||
cell.pos = p0 - 1;
|
||||
return true;
|
||||
}
|
||||
return false;
|
||||
}
|
||||
// invalidate tails which will be cleared
|
||||
@@ -368,6 +389,13 @@ llama_pos llama_memory_recurrent::seq_pos_max(llama_seq_id seq_id) const {
|
||||
return result;
|
||||
}
|
||||
|
||||
void llama_memory_recurrent::set_rs_idx(llama_seq_id seq_id, uint32_t idx) {
|
||||
if (seq_id < 0 || (size_t) seq_id >= rs_idx.size()) {
|
||||
return;
|
||||
}
|
||||
rs_idx[seq_id] = (idx > n_rs_seq) ? n_rs_seq : idx;
|
||||
}
|
||||
|
||||
std::map<ggml_backend_buffer_type_t, size_t> llama_memory_recurrent::memory_breakdown() const {
|
||||
std::map<ggml_backend_buffer_type_t, size_t> ret;
|
||||
for (const auto & [_, buf] : ctxs_bufs) {
|
||||
@@ -703,6 +731,7 @@ void llama_memory_recurrent::state_write(llama_io_write_i & io, llama_seq_id seq
|
||||
GGML_UNUSED(flags);
|
||||
|
||||
std::vector<std::pair<uint32_t, uint32_t>> cell_ranges; // ranges, from inclusive, to exclusive
|
||||
std::vector<std::pair<uint32_t, uint32_t>> cell_ranges_data; // logical source row ranges
|
||||
uint32_t cell_count = 0;
|
||||
|
||||
// Count the number of cells with the specified seq_id
|
||||
@@ -712,6 +741,35 @@ void llama_memory_recurrent::state_write(llama_io_write_i & io, llama_seq_id seq
|
||||
const auto & cell = cells[i];
|
||||
if ((seq_id == -1 && !cell.is_empty()) || cell.has_seq_id(seq_id)) {
|
||||
++cell_count;
|
||||
uint32_t rs_idx_cur = 0;
|
||||
|
||||
if (n_rs_seq != 0) {
|
||||
if (seq_id != -1) {
|
||||
GGML_ASSERT(seq_id >= 0 && (size_t) seq_id < rs_idx.size());
|
||||
rs_idx_cur = rs_idx[seq_id];
|
||||
} else {
|
||||
bool has_rs_idx = false;
|
||||
for (const llama_seq_id cell_seq_id : cell.seq_id) {
|
||||
GGML_ASSERT(cell_seq_id >= 0 && (size_t) cell_seq_id < rs_idx.size());
|
||||
|
||||
const uint32_t seq_rs_idx = rs_idx[cell_seq_id];
|
||||
if (!has_rs_idx) {
|
||||
rs_idx_cur = seq_rs_idx;
|
||||
has_rs_idx = true;
|
||||
} else if (rs_idx_cur != seq_rs_idx) {
|
||||
GGML_ABORT("cannot write shared recurrent state with different rollback indices");
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
const uint32_t cell_id = rs_idx_cur * size + (cell.src >= 0 ? cell.src : (int32_t) i);
|
||||
if (cell_ranges_data.empty() || cell_ranges_data.back().second != cell_id) {
|
||||
cell_ranges_data.emplace_back(cell_id, cell_id + 1);
|
||||
} else {
|
||||
cell_ranges_data.back().second++;
|
||||
}
|
||||
|
||||
if (cell_range_begin == size) {
|
||||
cell_range_begin = i;
|
||||
}
|
||||
@@ -726,7 +784,7 @@ void llama_memory_recurrent::state_write(llama_io_write_i & io, llama_seq_id seq
|
||||
cell_ranges.emplace_back(cell_range_begin, size);
|
||||
}
|
||||
|
||||
if (flags % LLAMA_STATE_SEQ_FLAGS_ON_DEVICE && cell_ranges.size() > 1) {
|
||||
if ((flags & LLAMA_STATE_SEQ_FLAGS_ON_DEVICE) && cell_ranges.size() > 1) {
|
||||
GGML_ABORT("cannot save/load multiple ranges of cells to/from device memory\n");
|
||||
}
|
||||
|
||||
@@ -737,10 +795,16 @@ void llama_memory_recurrent::state_write(llama_io_write_i & io, llama_seq_id seq
|
||||
}
|
||||
GGML_ASSERT(cell_count == cell_count_check);
|
||||
|
||||
cell_count_check = 0;
|
||||
for (const auto & range : cell_ranges_data) {
|
||||
cell_count_check += range.second - range.first;
|
||||
}
|
||||
GGML_ASSERT(cell_count == cell_count_check);
|
||||
|
||||
io.write(&cell_count, sizeof(cell_count));
|
||||
|
||||
state_write_meta(io, cell_ranges, seq_id);
|
||||
state_write_data(io, cell_ranges);
|
||||
state_write_data(io, cell_ranges_data);
|
||||
}
|
||||
|
||||
void llama_memory_recurrent::state_read(llama_io_read_i & io, llama_seq_id seq_id, llama_state_seq_flags flags) {
|
||||
@@ -762,6 +826,14 @@ void llama_memory_recurrent::state_read(llama_io_read_i & io, llama_seq_id seq_i
|
||||
}
|
||||
throw std::runtime_error("failed to restore kv cache");
|
||||
}
|
||||
|
||||
if (n_rs_seq != 0) {
|
||||
if (seq_id == -1) {
|
||||
std::fill(rs_idx.begin(), rs_idx.end(), 0);
|
||||
} else {
|
||||
set_rs_idx(seq_id, 0);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
void llama_memory_recurrent::state_write_meta(llama_io_write_i & io, const std::vector<std::pair<uint32_t, uint32_t>> & cell_ranges, llama_seq_id seq_id) const {
|
||||
@@ -804,7 +876,8 @@ void llama_memory_recurrent::state_write_data(llama_io_write_i & io, const std::
|
||||
const uint64_t r_size_row = ggml_row_size(r_l[il]->type, hparams.n_embd_r());
|
||||
io.write(&r_size_row, sizeof(r_size_row));
|
||||
|
||||
// Write each range of cells of r_size_row length
|
||||
// Write each logical cell row range. With pending recurrent rollback,
|
||||
// the logical current state may live in a rollback snapshot plane.
|
||||
for (const auto & range : cell_ranges) {
|
||||
const size_t range_size = range.second - range.first;
|
||||
const size_t buf_size = range_size * r_size_row;
|
||||
@@ -825,7 +898,8 @@ void llama_memory_recurrent::state_write_data(llama_io_write_i & io, const std::
|
||||
const uint64_t s_size_row = ggml_row_size(s_l[il]->type, hparams.n_embd_s());
|
||||
io.write(&s_size_row, sizeof(s_size_row));
|
||||
|
||||
// Write each range of S tensor rows
|
||||
// Write each logical cell row range. With pending recurrent rollback,
|
||||
// the logical current state may live in a rollback snapshot plane.
|
||||
for (const auto & range : cell_ranges) {
|
||||
const size_t range_size = range.second - range.first;
|
||||
const size_t buf_size = range_size * s_size_row;
|
||||
@@ -852,9 +926,8 @@ void llama_memory_recurrent::state_write_data(llama_io_write_i & io, const std::
|
||||
// Write GQA embedding size
|
||||
io.write(&n_embd_s, sizeof(n_embd_s));
|
||||
|
||||
// For each row, we get the element values of each cell
|
||||
// For each row, we get the element values of each logical cell
|
||||
for (uint32_t j = 0; j < n_embd_s; ++j) {
|
||||
// Write each range of cells of s_size_el length
|
||||
for (const auto & range : cell_ranges) {
|
||||
const size_t range_size = range.second - range.first;
|
||||
const size_t src_offset = (range.first + j * mem_size) * s_size_el;
|
||||
@@ -1163,5 +1236,21 @@ ggml_tensor * llama_memory_recurrent_context::get_s_l(int32_t il) const {
|
||||
}
|
||||
|
||||
int32_t llama_memory_recurrent_context::s_copy(int i) const {
|
||||
return mem->cells[i + mem->head].src0;
|
||||
const uint32_t cell_idx = i + mem->head;
|
||||
const int32_t src0 = mem->cells[cell_idx].src0;
|
||||
|
||||
if (mem->n_rs_seq == 0) {
|
||||
return src0;
|
||||
}
|
||||
|
||||
uint32_t idx = 0;
|
||||
if (!mem->cells[cell_idx].seq_id.empty()) {
|
||||
const llama_seq_id seq = *mem->cells[cell_idx].seq_id.begin();
|
||||
if (seq >= 0 && (size_t) seq < mem->rs_idx.size()) {
|
||||
idx = mem->rs_idx[seq];
|
||||
// reset rollback idx
|
||||
mem->rs_idx[seq] = 0;
|
||||
}
|
||||
}
|
||||
return (int32_t)(idx * mem->size) + src0;
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user