DeepseekV4: fix seq_rm (#25588)
* DeepseekV4: fix seq_rm * implement proper seq_cp * create actual update context
This commit is contained in:
+85
-15
@@ -720,7 +720,7 @@ llama_dsv4_comp_state::llama_dsv4_comp_state(
|
|||||||
auto it = ctx_map.find(buft);
|
auto it = ctx_map.find(buft);
|
||||||
if (it == ctx_map.end()) {
|
if (it == ctx_map.end()) {
|
||||||
ggml_init_params params = {
|
ggml_init_params params = {
|
||||||
/*.mem_size =*/ size_t(2u*hparams.n_layer()*ggml_tensor_overhead()),
|
/*.mem_size =*/ size_t(2u*(1 + n_stream)*hparams.n_layer()*ggml_tensor_overhead()),
|
||||||
/*.mem_buffer =*/ NULL,
|
/*.mem_buffer =*/ NULL,
|
||||||
/*.no_alloc =*/ true,
|
/*.no_alloc =*/ true,
|
||||||
};
|
};
|
||||||
@@ -767,9 +767,17 @@ llama_dsv4_comp_state::llama_dsv4_comp_state(
|
|||||||
ggml_format_name(kv, "dsv4_%s_state_kv_l%d", name, il);
|
ggml_format_name(kv, "dsv4_%s_state_kv_l%d", name, il);
|
||||||
ggml_format_name(score, "dsv4_%s_state_score_l%d", name, il);
|
ggml_format_name(score, "dsv4_%s_state_score_l%d", name, il);
|
||||||
|
|
||||||
|
std::vector<ggml_tensor *> kv_stream;
|
||||||
|
std::vector<ggml_tensor *> score_stream;
|
||||||
|
|
||||||
|
for (uint32_t s = 0; s < n_stream; ++s) {
|
||||||
|
kv_stream.push_back(ggml_view_2d(ctx, kv, n_embd_state, state_size, kv->nb[1], s*kv->nb[2]));
|
||||||
|
score_stream.push_back(ggml_view_2d(ctx, score, n_embd_state, state_size, score->nb[1], s*score->nb[2]));
|
||||||
|
}
|
||||||
|
|
||||||
map_layer_ids[il] = layers.size();
|
map_layer_ids[il] = layers.size();
|
||||||
|
|
||||||
layers.push_back({ il, kv, score });
|
layers.push_back({ il, kv, score, std::move(kv_stream), std::move(score_stream) });
|
||||||
}
|
}
|
||||||
|
|
||||||
for (auto & [buft, ctx] : ctx_map) {
|
for (auto & [buft, ctx] : ctx_map) {
|
||||||
@@ -809,6 +817,30 @@ void llama_dsv4_comp_state::clear(llama_seq_id seq_id, bool data) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
void llama_dsv4_comp_state::seq_cp(llama_seq_id seq_id_src, llama_seq_id seq_id_dst) {
|
||||||
|
GGML_ASSERT(seq_id_src >= 0 && (uint32_t) seq_id_src < n_stream);
|
||||||
|
GGML_ASSERT(seq_id_dst >= 0 && (uint32_t) seq_id_dst < n_stream);
|
||||||
|
|
||||||
|
if (seq_id_src == seq_id_dst) {
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
|
||||||
|
sc_info.ssrc.push_back((uint32_t) seq_id_src);
|
||||||
|
sc_info.sdst.push_back((uint32_t) seq_id_dst);
|
||||||
|
}
|
||||||
|
|
||||||
|
void llama_dsv4_comp_state::apply_copies(const stream_copy_info & sc_info) const {
|
||||||
|
for (size_t i = 0; i < sc_info.ssrc.size(); ++i) {
|
||||||
|
const uint32_t ssrc = sc_info.ssrc[i];
|
||||||
|
const uint32_t sdst = sc_info.sdst[i];
|
||||||
|
|
||||||
|
for (const auto & layer : layers) {
|
||||||
|
ggml_backend_tensor_copy(layer.kv_stream[ssrc], layer.kv_stream[sdst]);
|
||||||
|
ggml_backend_tensor_copy(layer.score_stream[ssrc], layer.score_stream[sdst]);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
uint32_t llama_dsv4_comp_state::get_ratio() const {
|
uint32_t llama_dsv4_comp_state::get_ratio() const {
|
||||||
return ratio;
|
return ratio;
|
||||||
}
|
}
|
||||||
@@ -1154,7 +1186,13 @@ llama_memory_context_ptr llama_kv_cache_dsv4::init_full() {
|
|||||||
}
|
}
|
||||||
|
|
||||||
llama_memory_context_ptr llama_kv_cache_dsv4::init_update(llama_context * lctx, bool optimize) {
|
llama_memory_context_ptr llama_kv_cache_dsv4::init_update(llama_context * lctx, bool optimize) {
|
||||||
return std::make_unique<llama_kv_cache_dsv4_context>(this, lctx, optimize);
|
return std::make_unique<llama_kv_cache_dsv4_context>(
|
||||||
|
this,
|
||||||
|
lctx,
|
||||||
|
optimize,
|
||||||
|
std::move(csa_state->sc_info),
|
||||||
|
std::move(hca_state->sc_info),
|
||||||
|
std::move(lid_state->sc_info));
|
||||||
}
|
}
|
||||||
|
|
||||||
bool llama_kv_cache_dsv4::get_can_shift() const {
|
bool llama_kv_cache_dsv4::get_can_shift() const {
|
||||||
@@ -1174,14 +1212,19 @@ bool llama_kv_cache_dsv4::seq_rm(llama_seq_id seq_id, llama_pos p0, llama_pos p1
|
|||||||
}
|
}
|
||||||
|
|
||||||
if (p0 > 0) {
|
if (p0 > 0) {
|
||||||
// DSV4 compressed cache rows are derived from running compressor state,
|
if (seq_id < 0 || (uint32_t) seq_id >= n_seq_max ||
|
||||||
// so arbitrary rollback is not reconstructible from the raw cache alone.
|
p0 <= kv_raw->seq_pos_max(seq_id)) {
|
||||||
// Allow the common prompt-cache cleanup no-op: remove [end, infinity).
|
return false;
|
||||||
if (seq_id >= 0 && p0 > kv_raw->seq_pos_max(seq_id)) {
|
|
||||||
return true;
|
|
||||||
}
|
}
|
||||||
|
|
||||||
return false;
|
bool res = true;
|
||||||
|
|
||||||
|
res = res & kv_raw->seq_rm(seq_id, p0, -1);
|
||||||
|
res = res & kv_csa->seq_rm(seq_id, p0/DSV4_CSA_RATIO, -1);
|
||||||
|
res = res & kv_hca->seq_rm(seq_id, p0/DSV4_HCA_RATIO, -1);
|
||||||
|
res = res & kv_lid->seq_rm(seq_id, p0/DSV4_CSA_RATIO, -1);
|
||||||
|
|
||||||
|
return res;
|
||||||
}
|
}
|
||||||
|
|
||||||
const bool res = kv_raw->seq_rm(seq_id, p0, p1);
|
const bool res = kv_raw->seq_rm(seq_id, p0, p1);
|
||||||
@@ -1194,7 +1237,16 @@ bool llama_kv_cache_dsv4::seq_rm(llama_seq_id seq_id, llama_pos p0, llama_pos p1
|
|||||||
}
|
}
|
||||||
|
|
||||||
void llama_kv_cache_dsv4::seq_cp(llama_seq_id seq_id_src, llama_seq_id seq_id_dst, llama_pos p0, llama_pos p1) {
|
void llama_kv_cache_dsv4::seq_cp(llama_seq_id seq_id_src, llama_seq_id seq_id_dst, llama_pos p0, llama_pos p1) {
|
||||||
|
GGML_ASSERT(p0 <= 0 && p1 < 0 && "DSV4 only supports full sequence copies");
|
||||||
|
|
||||||
kv_raw->seq_cp(seq_id_src, seq_id_dst, p0, p1);
|
kv_raw->seq_cp(seq_id_src, seq_id_dst, p0, p1);
|
||||||
|
kv_csa->seq_cp(seq_id_src, seq_id_dst, -1, -1);
|
||||||
|
kv_hca->seq_cp(seq_id_src, seq_id_dst, -1, -1);
|
||||||
|
kv_lid->seq_cp(seq_id_src, seq_id_dst, -1, -1);
|
||||||
|
|
||||||
|
csa_state->seq_cp(seq_id_src, seq_id_dst);
|
||||||
|
hca_state->seq_cp(seq_id_src, seq_id_dst);
|
||||||
|
lid_state->seq_cp(seq_id_src, seq_id_dst);
|
||||||
}
|
}
|
||||||
|
|
||||||
void llama_kv_cache_dsv4::seq_keep(llama_seq_id seq_id) {
|
void llama_kv_cache_dsv4::seq_keep(llama_seq_id seq_id) {
|
||||||
@@ -1639,20 +1691,26 @@ llama_kv_cache_dsv4_context::llama_kv_cache_dsv4_context(
|
|||||||
llama_kv_cache_dsv4_context::llama_kv_cache_dsv4_context(
|
llama_kv_cache_dsv4_context::llama_kv_cache_dsv4_context(
|
||||||
llama_kv_cache_dsv4 * kv,
|
llama_kv_cache_dsv4 * kv,
|
||||||
llama_context * lctx,
|
llama_context * lctx,
|
||||||
bool optimize) :
|
bool optimize,
|
||||||
|
stream_copy_info sc_info_csa,
|
||||||
|
stream_copy_info sc_info_hca,
|
||||||
|
stream_copy_info sc_info_lid) :
|
||||||
ctx_raw(std::make_unique<llama_kv_cache_dsv4_raw_context>(kv->get_raw(), lctx, optimize)),
|
ctx_raw(std::make_unique<llama_kv_cache_dsv4_raw_context>(kv->get_raw(), lctx, optimize)),
|
||||||
ctx_csa_mem(kv->get_csa()->init_update(lctx, optimize)),
|
ctx_csa_mem(kv->get_csa()->init_update(lctx, optimize)),
|
||||||
ctx_hca_mem(kv->get_hca()->init_update(lctx, optimize)),
|
ctx_hca_mem(kv->get_hca()->init_update(lctx, optimize)),
|
||||||
ctx_lid_mem(kv->get_lid()->init_update(lctx, optimize)),
|
ctx_lid_mem(kv->get_lid()->init_update(lctx, optimize)),
|
||||||
ctx_csa(std::make_unique<llama_kv_cache_dsv4_comp_context>(kv->get_csa())),
|
|
||||||
ctx_hca(std::make_unique<llama_kv_cache_dsv4_comp_context>(kv->get_hca())),
|
|
||||||
ctx_lid(std::make_unique<llama_kv_cache_dsv4_comp_context>(kv->get_lid())),
|
|
||||||
csa_state(kv->get_csa_state()),
|
csa_state(kv->get_csa_state()),
|
||||||
hca_state(kv->get_hca_state()),
|
hca_state(kv->get_hca_state()),
|
||||||
lid_state(kv->get_lid_state()),
|
lid_state(kv->get_lid_state()),
|
||||||
|
sc_info_csa(std::move(sc_info_csa)),
|
||||||
|
sc_info_hca(std::move(sc_info_hca)),
|
||||||
|
sc_info_lid(std::move(sc_info_lid)),
|
||||||
status(llama_memory_status_combine(
|
status(llama_memory_status_combine(
|
||||||
llama_memory_status_combine(ctx_raw->get_status(), ctx_csa_mem->get_status()),
|
llama_memory_status_combine(
|
||||||
llama_memory_status_combine(ctx_hca_mem->get_status(), ctx_lid_mem->get_status()))) {
|
llama_memory_status_combine(ctx_raw->get_status(), ctx_csa_mem->get_status()),
|
||||||
|
llama_memory_status_combine(ctx_hca_mem->get_status(), ctx_lid_mem->get_status())),
|
||||||
|
this->sc_info_csa.empty() && this->sc_info_hca.empty() && this->sc_info_lid.empty() ?
|
||||||
|
LLAMA_MEMORY_STATUS_NO_UPDATE : LLAMA_MEMORY_STATUS_SUCCESS)) {
|
||||||
}
|
}
|
||||||
|
|
||||||
llama_kv_cache_dsv4_context::llama_kv_cache_dsv4_context(
|
llama_kv_cache_dsv4_context::llama_kv_cache_dsv4_context(
|
||||||
@@ -1720,6 +1778,18 @@ bool llama_kv_cache_dsv4_context::apply() {
|
|||||||
|
|
||||||
res = res & ctx_raw->apply();
|
res = res & ctx_raw->apply();
|
||||||
|
|
||||||
|
if (ctx_csa_mem) {
|
||||||
|
res = res & ctx_csa_mem->apply();
|
||||||
|
res = res & ctx_hca_mem->apply();
|
||||||
|
res = res & ctx_lid_mem->apply();
|
||||||
|
}
|
||||||
|
|
||||||
|
if (ubatches.empty()) {
|
||||||
|
csa_state->apply_copies(sc_info_csa);
|
||||||
|
hca_state->apply_copies(sc_info_hca);
|
||||||
|
lid_state->apply_copies(sc_info_lid);
|
||||||
|
}
|
||||||
|
|
||||||
return res;
|
return res;
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -10,6 +10,10 @@
|
|||||||
|
|
||||||
class llama_dsv4_comp_state {
|
class llama_dsv4_comp_state {
|
||||||
public:
|
public:
|
||||||
|
using stream_copy_info = llama_kv_cache::stream_copy_info;
|
||||||
|
|
||||||
|
stream_copy_info sc_info;
|
||||||
|
|
||||||
llama_dsv4_comp_state(
|
llama_dsv4_comp_state(
|
||||||
const llama_model & model,
|
const llama_model & model,
|
||||||
bool offload,
|
bool offload,
|
||||||
@@ -22,6 +26,8 @@ public:
|
|||||||
const llama_memory_i::layer_filter_cb & filter);
|
const llama_memory_i::layer_filter_cb & filter);
|
||||||
|
|
||||||
void clear(llama_seq_id seq_id, bool data);
|
void clear(llama_seq_id seq_id, bool data);
|
||||||
|
void seq_cp(llama_seq_id seq_id_src, llama_seq_id seq_id_dst);
|
||||||
|
void apply_copies(const stream_copy_info & sc_info) const;
|
||||||
|
|
||||||
uint32_t get_ratio() const;
|
uint32_t get_ratio() const;
|
||||||
uint32_t get_state_size() const;
|
uint32_t get_state_size() const;
|
||||||
@@ -44,6 +50,9 @@ private:
|
|||||||
|
|
||||||
ggml_tensor * kv;
|
ggml_tensor * kv;
|
||||||
ggml_tensor * score;
|
ggml_tensor * score;
|
||||||
|
|
||||||
|
std::vector<ggml_tensor *> kv_stream;
|
||||||
|
std::vector<ggml_tensor *> score_stream;
|
||||||
};
|
};
|
||||||
|
|
||||||
const uint32_t ratio;
|
const uint32_t ratio;
|
||||||
@@ -245,6 +254,7 @@ private:
|
|||||||
class llama_kv_cache_dsv4_context : public llama_memory_context_i {
|
class llama_kv_cache_dsv4_context : public llama_memory_context_i {
|
||||||
public:
|
public:
|
||||||
using slot_info_vec_t = llama_kv_cache::slot_info_vec_t;
|
using slot_info_vec_t = llama_kv_cache::slot_info_vec_t;
|
||||||
|
using stream_copy_info = llama_kv_cache::stream_copy_info;
|
||||||
|
|
||||||
struct comp_plan {
|
struct comp_plan {
|
||||||
// Per-ubatch recipe for updating compressor state, committing completed
|
// Per-ubatch recipe for updating compressor state, committing completed
|
||||||
@@ -291,7 +301,10 @@ public:
|
|||||||
llama_kv_cache_dsv4_context(
|
llama_kv_cache_dsv4_context(
|
||||||
llama_kv_cache_dsv4 * kv,
|
llama_kv_cache_dsv4 * kv,
|
||||||
llama_context * lctx,
|
llama_context * lctx,
|
||||||
bool optimize);
|
bool optimize,
|
||||||
|
stream_copy_info sc_info_csa,
|
||||||
|
stream_copy_info sc_info_hca,
|
||||||
|
stream_copy_info sc_info_lid);
|
||||||
|
|
||||||
llama_kv_cache_dsv4_context(
|
llama_kv_cache_dsv4_context(
|
||||||
llama_kv_cache_dsv4 * kv,
|
llama_kv_cache_dsv4 * kv,
|
||||||
@@ -351,9 +364,13 @@ private:
|
|||||||
const std::unique_ptr<llama_kv_cache_dsv4_comp_context> ctx_hca;
|
const std::unique_ptr<llama_kv_cache_dsv4_comp_context> ctx_hca;
|
||||||
const std::unique_ptr<llama_kv_cache_dsv4_comp_context> ctx_lid;
|
const std::unique_ptr<llama_kv_cache_dsv4_comp_context> ctx_lid;
|
||||||
|
|
||||||
const llama_dsv4_comp_state * csa_state = nullptr;
|
llama_dsv4_comp_state * csa_state = nullptr;
|
||||||
const llama_dsv4_comp_state * hca_state = nullptr;
|
llama_dsv4_comp_state * hca_state = nullptr;
|
||||||
const llama_dsv4_comp_state * lid_state = nullptr;
|
llama_dsv4_comp_state * lid_state = nullptr;
|
||||||
|
|
||||||
|
stream_copy_info sc_info_csa;
|
||||||
|
stream_copy_info sc_info_hca;
|
||||||
|
stream_copy_info sc_info_lid;
|
||||||
|
|
||||||
bool reserve_plans = false;
|
bool reserve_plans = false;
|
||||||
mutable comp_plan reserve_plan_csa;
|
mutable comp_plan reserve_plan_csa;
|
||||||
|
|||||||
Reference in New Issue
Block a user