model: M3: Move MSA into a new memory implementation (#26338)
* Move MSA logic from llama-kv-cache into llama-kv-cache-msa * cont : minor * cont : ws fix --------- Co-authored-by: Georgi Gerganov <ggerganov@gmail.com>
This commit is contained in:
co-authored by
Georgi Gerganov
parent
563dec81c1
commit
67d5978bb1
@@ -25,6 +25,7 @@ add_library(llama
|
|||||||
llama-kv-cache.cpp
|
llama-kv-cache.cpp
|
||||||
llama-kv-cache-iswa.cpp
|
llama-kv-cache-iswa.cpp
|
||||||
llama-kv-cache-dsa.cpp
|
llama-kv-cache-dsa.cpp
|
||||||
|
llama-kv-cache-msa.cpp
|
||||||
llama-kv-cache-dsv4.cpp
|
llama-kv-cache-dsv4.cpp
|
||||||
llama-memory.cpp
|
llama-memory.cpp
|
||||||
llama-memory-hybrid.cpp
|
llama-memory-hybrid.cpp
|
||||||
|
|||||||
@@ -8,6 +8,7 @@
|
|||||||
#include "llama-kv-cache.h"
|
#include "llama-kv-cache.h"
|
||||||
#include "llama-kv-cache-iswa.h"
|
#include "llama-kv-cache-iswa.h"
|
||||||
#include "llama-kv-cache-dsa.h"
|
#include "llama-kv-cache-dsa.h"
|
||||||
|
#include "llama-kv-cache-msa.h"
|
||||||
#include "llama-kv-cache-dsv4.h"
|
#include "llama-kv-cache-dsv4.h"
|
||||||
#include "llama-memory-hybrid.h"
|
#include "llama-memory-hybrid.h"
|
||||||
#include "llama-memory-hybrid-iswa.h"
|
#include "llama-memory-hybrid-iswa.h"
|
||||||
@@ -518,6 +519,36 @@ bool llm_graph_input_attn_k::can_reuse(const llm_graph_params & params) {
|
|||||||
return res;
|
return res;
|
||||||
}
|
}
|
||||||
|
|
||||||
|
llm_graph_input_attn_kv_msa::llm_graph_input_attn_kv_msa(
|
||||||
|
const llama_hparams & hparams,
|
||||||
|
const llama_cparams & cparams,
|
||||||
|
const llama_kv_cache_msa_context * mctx) :
|
||||||
|
llm_graph_input_attn_kv(hparams, cparams, mctx->get_base()),
|
||||||
|
mctx_msa(mctx) {
|
||||||
|
}
|
||||||
|
|
||||||
|
void llm_graph_input_attn_kv_msa::set_input(const llama_ubatch * ubatch) {
|
||||||
|
llm_graph_input_attn_kv::set_input(ubatch);
|
||||||
|
|
||||||
|
mctx_msa->get_idx()->set_input_k_idxs(self_k_idxs_idx, ubatch);
|
||||||
|
}
|
||||||
|
|
||||||
|
bool llm_graph_input_attn_kv_msa::can_reuse(const llm_graph_params & params) {
|
||||||
|
mctx_msa = static_cast<const llama_kv_cache_msa_context *>(params.mctx);
|
||||||
|
|
||||||
|
// the parent class operates on the base cache context
|
||||||
|
this->mctx = mctx_msa->get_base();
|
||||||
|
|
||||||
|
bool res = true;
|
||||||
|
|
||||||
|
res &= self_k_idxs ->ne[0] == params.ubatch.n_tokens;
|
||||||
|
res &= self_k_idxs_idx->ne[0] == params.ubatch.n_tokens;
|
||||||
|
|
||||||
|
res &= can_reuse_kq_mask(self_kq_mask, this->mctx, params.ubatch, params.cparams);
|
||||||
|
|
||||||
|
return res;
|
||||||
|
}
|
||||||
|
|
||||||
void llm_graph_input_attn_k_dsa::set_input(const llama_ubatch * ubatch) {
|
void llm_graph_input_attn_k_dsa::set_input(const llama_ubatch * ubatch) {
|
||||||
mctx->get_mla()->set_input_k_idxs(self_k_idxs_mla, ubatch);
|
mctx->get_mla()->set_input_k_idxs(self_k_idxs_mla, ubatch);
|
||||||
|
|
||||||
@@ -3187,6 +3218,32 @@ llm_graph_input_attn_k_dsa * llm_graph_context::build_attn_inp_k_dsa() const {
|
|||||||
return (llm_graph_input_attn_k_dsa *) res->add_input(std::move(inp));
|
return (llm_graph_input_attn_k_dsa *) res->add_input(std::move(inp));
|
||||||
}
|
}
|
||||||
|
|
||||||
|
llm_graph_input_attn_kv_msa * llm_graph_context::build_attn_inp_kv_msa() const {
|
||||||
|
const auto * mctx_cur = static_cast<const llama_kv_cache_msa_context *>(mctx);
|
||||||
|
|
||||||
|
auto inp = std::make_unique<llm_graph_input_attn_kv_msa>(hparams, cparams, mctx_cur);
|
||||||
|
|
||||||
|
const auto * mctx_base = mctx_cur->get_base();
|
||||||
|
const auto * mctx_idx = mctx_cur->get_idx();
|
||||||
|
|
||||||
|
{
|
||||||
|
GGML_ASSERT(hparams.swa_type == LLAMA_SWA_TYPE_NONE && "Use llama_kv_cache_iswa for SWA");
|
||||||
|
|
||||||
|
inp->self_k_idxs = mctx_base->build_input_k_idxs(ctx0, ubatch);
|
||||||
|
inp->self_v_idxs = mctx_base->build_input_v_idxs(ctx0, ubatch);
|
||||||
|
|
||||||
|
inp->self_kq_mask = build_attn_inp_kq_mask(ctx0, mctx_base, ubatch, cparams);
|
||||||
|
inp->self_kq_mask_cnv = inp->self_kq_mask;
|
||||||
|
}
|
||||||
|
|
||||||
|
inp->self_k_rot = mctx_base->build_input_k_rot(ctx0);
|
||||||
|
inp->self_v_rot = mctx_base->build_input_v_rot(ctx0);
|
||||||
|
|
||||||
|
inp->self_k_idxs_idx = mctx_idx->build_input_k_idxs(ctx0, ubatch);
|
||||||
|
|
||||||
|
return (llm_graph_input_attn_kv_msa *) res->add_input(std::move(inp));
|
||||||
|
}
|
||||||
|
|
||||||
// TODO: maybe separate the inner implementation into a separate function
|
// TODO: maybe separate the inner implementation into a separate function
|
||||||
// like with the non-sliding window equivalent
|
// like with the non-sliding window equivalent
|
||||||
// once sliding-window hybrid caches are a thing.
|
// once sliding-window hybrid caches are a thing.
|
||||||
|
|||||||
@@ -23,6 +23,7 @@ struct llama_memory_context_i;
|
|||||||
|
|
||||||
class llama_kv_cache_context;
|
class llama_kv_cache_context;
|
||||||
class llama_kv_cache_dsa_context;
|
class llama_kv_cache_dsa_context;
|
||||||
|
class llama_kv_cache_msa_context;
|
||||||
class llama_kv_cache_dsv4_raw_context;
|
class llama_kv_cache_dsv4_raw_context;
|
||||||
class llama_kv_cache_dsv4_context;
|
class llama_kv_cache_dsv4_context;
|
||||||
class llama_kv_cache_iswa_context;
|
class llama_kv_cache_iswa_context;
|
||||||
@@ -425,6 +426,26 @@ public:
|
|||||||
const llama_kv_cache_dsa_context * mctx;
|
const llama_kv_cache_dsa_context * mctx;
|
||||||
};
|
};
|
||||||
|
|
||||||
|
// standard K/V attention input against the base cache, plus destination indices for the indexer key cache
|
||||||
|
class llm_graph_input_attn_kv_msa : public llm_graph_input_attn_kv {
|
||||||
|
public:
|
||||||
|
llm_graph_input_attn_kv_msa(
|
||||||
|
const llama_hparams & hparams,
|
||||||
|
const llama_cparams & cparams,
|
||||||
|
const llama_kv_cache_msa_context * mctx);
|
||||||
|
~llm_graph_input_attn_kv_msa() = default;
|
||||||
|
|
||||||
|
void set_input(const llama_ubatch * ubatch) override;
|
||||||
|
|
||||||
|
bool can_reuse(const llm_graph_params & params) override;
|
||||||
|
|
||||||
|
ggml_tensor * get_k_idxs_idx() const { return self_k_idxs_idx; }
|
||||||
|
|
||||||
|
ggml_tensor * self_k_idxs_idx = nullptr; // I64 [n_batch]
|
||||||
|
|
||||||
|
const llama_kv_cache_msa_context * mctx_msa;
|
||||||
|
};
|
||||||
|
|
||||||
class llm_graph_input_attn_kv_iswa : public llm_graph_input_i {
|
class llm_graph_input_attn_kv_iswa : public llm_graph_input_i {
|
||||||
public:
|
public:
|
||||||
llm_graph_input_attn_kv_iswa(
|
llm_graph_input_attn_kv_iswa(
|
||||||
@@ -1169,6 +1190,8 @@ struct llm_graph_context {
|
|||||||
|
|
||||||
llm_graph_input_attn_k_dsa * build_attn_inp_k_dsa() const;
|
llm_graph_input_attn_k_dsa * build_attn_inp_k_dsa() const;
|
||||||
|
|
||||||
|
llm_graph_input_attn_kv_msa * build_attn_inp_kv_msa() const;
|
||||||
|
|
||||||
ggml_tensor * build_attn(
|
ggml_tensor * build_attn(
|
||||||
llm_graph_input_attn_k_dsa * inp,
|
llm_graph_input_attn_k_dsa * inp,
|
||||||
ggml_tensor * wo,
|
ggml_tensor * wo,
|
||||||
|
|||||||
@@ -180,16 +180,6 @@ uint32_t llama_hparams::n_embd_v_gqa_max() const {
|
|||||||
return val;
|
return val;
|
||||||
}
|
}
|
||||||
|
|
||||||
uint32_t llama_hparams::n_embd_k_idx(uint32_t il) const {
|
|
||||||
if (!indexer_kv || indexer_head_size == 0) {
|
|
||||||
return 0; // arch without a MSA indexer
|
|
||||||
}
|
|
||||||
if (il < n_layer_dense_lead) {
|
|
||||||
return 0; // leading dense layers carry no indexer
|
|
||||||
}
|
|
||||||
return indexer_head_size; // 128
|
|
||||||
}
|
|
||||||
|
|
||||||
uint32_t llama_hparams::n_embd_r() const {
|
uint32_t llama_hparams::n_embd_r() const {
|
||||||
if (wkv_head_size != 0) {
|
if (wkv_head_size != 0) {
|
||||||
// for RWKV models
|
// for RWKV models
|
||||||
|
|||||||
@@ -230,8 +230,6 @@ struct llama_hparams {
|
|||||||
// MSA
|
// MSA
|
||||||
uint32_t indexer_block_size = 0;
|
uint32_t indexer_block_size = 0;
|
||||||
uint32_t indexer_local_blocks = 0;
|
uint32_t indexer_local_blocks = 0;
|
||||||
// MSA stores its indexer keys in the main KV cache (k_idx tensors);
|
|
||||||
bool indexer_kv = false;
|
|
||||||
|
|
||||||
// Indexer is "full" (1) or "shared" (0)
|
// Indexer is "full" (1) or "shared" (0)
|
||||||
// Shared indexers reuse top-k from previous full layer
|
// Shared indexers reuse top-k from previous full layer
|
||||||
@@ -356,9 +354,6 @@ struct llama_hparams {
|
|||||||
uint32_t n_embd_k_gqa_max() const;
|
uint32_t n_embd_k_gqa_max() const;
|
||||||
uint32_t n_embd_v_gqa_max() const;
|
uint32_t n_embd_v_gqa_max() const;
|
||||||
|
|
||||||
// dimension of the single-head MSA indexer key stream
|
|
||||||
uint32_t n_embd_k_idx(uint32_t il = 0) const;
|
|
||||||
|
|
||||||
// dimension of the rolling state embeddings
|
// dimension of the rolling state embeddings
|
||||||
// corresponds to Mamba's conv_states size or RWKV's token_shift states size
|
// corresponds to Mamba's conv_states size or RWKV's token_shift states size
|
||||||
uint32_t n_embd_r() const;
|
uint32_t n_embd_r() const;
|
||||||
|
|||||||
@@ -0,0 +1,395 @@
|
|||||||
|
#include "llama-kv-cache-msa.h"
|
||||||
|
|
||||||
|
#include "llama-impl.h"
|
||||||
|
#include "llama-batch.h"
|
||||||
|
#include "llama-model.h"
|
||||||
|
|
||||||
|
#include <algorithm>
|
||||||
|
#include <cassert>
|
||||||
|
#include <cmath>
|
||||||
|
|
||||||
|
// llama_kv_cache_msa
|
||||||
|
|
||||||
|
llama_kv_cache_msa::llama_kv_cache_msa(
|
||||||
|
const llama_model & model,
|
||||||
|
ggml_type type_k,
|
||||||
|
ggml_type type_v,
|
||||||
|
bool v_trans,
|
||||||
|
bool offload,
|
||||||
|
bool unified,
|
||||||
|
uint32_t kv_size,
|
||||||
|
uint32_t n_seq_max,
|
||||||
|
uint32_t n_pad,
|
||||||
|
uint32_t n_swa,
|
||||||
|
llama_swa_type swa_type,
|
||||||
|
const layer_filter_cb & filter,
|
||||||
|
const layer_filter_cb & filter_idx,
|
||||||
|
const layer_reuse_cb & reuse) :
|
||||||
|
hparams_idx(model.hparams),
|
||||||
|
n_stream(unified ? 1 : n_seq_max), n_seq_max(n_seq_max), n_pad(n_pad),
|
||||||
|
n_swa(n_swa), swa_type(swa_type) {
|
||||||
|
|
||||||
|
LLAMA_LOG_INFO("%s: creating main KV cache, size = %u cells\n", __func__, kv_size);
|
||||||
|
|
||||||
|
kv_base = std::make_unique<llama_kv_cache>(
|
||||||
|
model, model.hparams, type_k, type_v,
|
||||||
|
v_trans, offload, unified, kv_size, n_seq_max, n_pad,
|
||||||
|
n_swa, swa_type, nullptr, filter, reuse, nullptr);
|
||||||
|
|
||||||
|
// the MSA indexer uses a single key head per layer
|
||||||
|
std::fill(hparams_idx.n_head_kv_arr.begin(), hparams_idx.n_head_kv_arr.end(), 1);
|
||||||
|
hparams_idx.n_embd_head_k_full = model.hparams.indexer_head_size;
|
||||||
|
// the rope parameters are kept identical to the main cache
|
||||||
|
|
||||||
|
LLAMA_LOG_INFO("%s: creating indexer KV cache, size = %u cells\n", __func__, kv_size);
|
||||||
|
|
||||||
|
kv_idx = std::make_unique<llama_kv_cache>(
|
||||||
|
model, hparams_idx, type_k, type_v,
|
||||||
|
v_trans, offload, unified, kv_size, n_seq_max, n_pad,
|
||||||
|
n_swa, swa_type, nullptr, filter_idx, reuse, nullptr);
|
||||||
|
}
|
||||||
|
|
||||||
|
void llama_kv_cache_msa::clear(bool data) {
|
||||||
|
kv_base->clear(data);
|
||||||
|
kv_idx ->clear(data);
|
||||||
|
}
|
||||||
|
|
||||||
|
bool llama_kv_cache_msa::seq_rm(llama_seq_id seq_id, llama_pos p0, llama_pos p1) {
|
||||||
|
bool res = true;
|
||||||
|
|
||||||
|
res = res & kv_base->seq_rm(seq_id, p0, p1);
|
||||||
|
res = res & kv_idx ->seq_rm(seq_id, p0, p1);
|
||||||
|
|
||||||
|
return res;
|
||||||
|
}
|
||||||
|
|
||||||
|
void llama_kv_cache_msa::seq_cp(llama_seq_id seq_id_src, llama_seq_id seq_id_dst, llama_pos p0, llama_pos p1) {
|
||||||
|
kv_base->seq_cp(seq_id_src, seq_id_dst, p0, p1);
|
||||||
|
kv_idx ->seq_cp(seq_id_src, seq_id_dst, p0, p1);
|
||||||
|
}
|
||||||
|
|
||||||
|
void llama_kv_cache_msa::seq_keep(llama_seq_id seq_id) {
|
||||||
|
kv_base->seq_keep(seq_id);
|
||||||
|
kv_idx ->seq_keep(seq_id);
|
||||||
|
}
|
||||||
|
|
||||||
|
void llama_kv_cache_msa::seq_add(llama_seq_id seq_id, llama_pos p0, llama_pos p1, llama_pos shift) {
|
||||||
|
kv_base->seq_add(seq_id, p0, p1, shift);
|
||||||
|
kv_idx ->seq_add(seq_id, p0, p1, shift);
|
||||||
|
}
|
||||||
|
|
||||||
|
void llama_kv_cache_msa::seq_div(llama_seq_id seq_id, llama_pos p0, llama_pos p1, int d) {
|
||||||
|
kv_base->seq_div(seq_id, p0, p1, d);
|
||||||
|
kv_idx ->seq_div(seq_id, p0, p1, d);
|
||||||
|
}
|
||||||
|
|
||||||
|
llama_pos llama_kv_cache_msa::seq_pos_min(llama_seq_id seq_id) const {
|
||||||
|
return kv_base->seq_pos_min(seq_id);
|
||||||
|
}
|
||||||
|
|
||||||
|
llama_pos llama_kv_cache_msa::seq_pos_max(llama_seq_id seq_id) const {
|
||||||
|
return kv_base->seq_pos_max(seq_id);
|
||||||
|
}
|
||||||
|
|
||||||
|
std::map<ggml_backend_buffer_type_t, size_t> llama_kv_cache_msa::memory_breakdown() const {
|
||||||
|
std::map<ggml_backend_buffer_type_t, size_t> mb = kv_base->memory_breakdown();
|
||||||
|
for (const auto & buft_size : kv_idx->memory_breakdown()) {
|
||||||
|
mb[buft_size.first] += buft_size.second;
|
||||||
|
}
|
||||||
|
return mb;
|
||||||
|
}
|
||||||
|
|
||||||
|
llama_memory_context_ptr llama_kv_cache_msa::init_batch(
|
||||||
|
llama_batch_allocr & balloc,
|
||||||
|
uint32_t n_ubatch,
|
||||||
|
bool embd_all) {
|
||||||
|
GGML_UNUSED(embd_all);
|
||||||
|
|
||||||
|
do {
|
||||||
|
balloc.split_reset();
|
||||||
|
|
||||||
|
std::vector<llama_ubatch> ubatches;
|
||||||
|
while (true) {
|
||||||
|
auto ubatch = n_stream == 1 ? balloc.split_simple(n_ubatch) : balloc.split_equal(n_ubatch, true, 0);
|
||||||
|
|
||||||
|
if (ubatch.n_tokens == 0) {
|
||||||
|
break;
|
||||||
|
}
|
||||||
|
|
||||||
|
ubatches.push_back(std::move(ubatch));
|
||||||
|
}
|
||||||
|
|
||||||
|
if (balloc.get_n_used() < balloc.get_n_tokens()) {
|
||||||
|
// failed to find a suitable split
|
||||||
|
break;
|
||||||
|
}
|
||||||
|
|
||||||
|
auto sinfos_base = kv_base->prepare(ubatches);
|
||||||
|
if (sinfos_base.empty()) {
|
||||||
|
break;
|
||||||
|
}
|
||||||
|
|
||||||
|
auto sinfos_idx = kv_idx->prepare(ubatches);
|
||||||
|
if (sinfos_idx.empty()) {
|
||||||
|
break;
|
||||||
|
}
|
||||||
|
|
||||||
|
assert(sinfos_base.size() == sinfos_idx.size());
|
||||||
|
|
||||||
|
return std::make_unique<llama_kv_cache_msa_context>(
|
||||||
|
this, std::move(sinfos_base), std::move(sinfos_idx), std::move(ubatches));
|
||||||
|
} while (false);
|
||||||
|
|
||||||
|
return std::make_unique<llama_kv_cache_msa_context>(LLAMA_MEMORY_STATUS_FAILED_PREPARE);
|
||||||
|
}
|
||||||
|
|
||||||
|
llama_memory_context_ptr llama_kv_cache_msa::init_full() {
|
||||||
|
return std::make_unique<llama_kv_cache_msa_context>(this);
|
||||||
|
}
|
||||||
|
|
||||||
|
llama_memory_context_ptr llama_kv_cache_msa::init_update(llama_context * lctx, bool optimize) {
|
||||||
|
return std::make_unique<llama_kv_cache_msa_context>(this, lctx, optimize);
|
||||||
|
}
|
||||||
|
|
||||||
|
bool llama_kv_cache_msa::get_can_shift() const {
|
||||||
|
return kv_base->get_can_shift() &&
|
||||||
|
kv_idx ->get_can_shift() &&
|
||||||
|
kv_base->get_size() == kv_idx->get_size();
|
||||||
|
}
|
||||||
|
|
||||||
|
void llama_kv_cache_msa::state_write(llama_io_write_i & io, llama_seq_id seq_id, llama_state_seq_flags flags) const {
|
||||||
|
kv_base->state_write(io, seq_id, flags);
|
||||||
|
kv_idx ->state_write(io, seq_id, flags);
|
||||||
|
}
|
||||||
|
|
||||||
|
void llama_kv_cache_msa::state_read(llama_io_read_i & io, llama_seq_id seq_id, llama_state_seq_flags flags) {
|
||||||
|
kv_base->state_read(io, seq_id, flags);
|
||||||
|
kv_idx ->state_read(io, seq_id, flags);
|
||||||
|
}
|
||||||
|
|
||||||
|
llama_kv_cache * llama_kv_cache_msa::get_base() const {
|
||||||
|
return kv_base.get();
|
||||||
|
}
|
||||||
|
|
||||||
|
llama_kv_cache * llama_kv_cache_msa::get_idx() const {
|
||||||
|
return kv_idx.get();
|
||||||
|
}
|
||||||
|
|
||||||
|
// llama_kv_cache_msa_context
|
||||||
|
|
||||||
|
llama_kv_cache_msa_context::llama_kv_cache_msa_context(llama_memory_status status) :
|
||||||
|
kv(nullptr), status(status) {}
|
||||||
|
|
||||||
|
llama_kv_cache_msa_context::llama_kv_cache_msa_context(
|
||||||
|
llama_kv_cache_msa * kv) :
|
||||||
|
kv(kv),
|
||||||
|
ctx_base(kv->get_base()->init_full()),
|
||||||
|
ctx_idx (kv->get_idx ()->init_full()),
|
||||||
|
status(llama_memory_status_combine(ctx_base->get_status(), ctx_idx->get_status())) {
|
||||||
|
}
|
||||||
|
|
||||||
|
llama_kv_cache_msa_context::llama_kv_cache_msa_context(
|
||||||
|
llama_kv_cache_msa * kv,
|
||||||
|
llama_context * lctx,
|
||||||
|
bool optimize) :
|
||||||
|
kv(kv),
|
||||||
|
ctx_base(kv->get_base()->init_update(lctx, optimize)),
|
||||||
|
ctx_idx (kv->get_idx ()->init_update(lctx, optimize)),
|
||||||
|
status(llama_memory_status_combine(ctx_base->get_status(), ctx_idx->get_status())) {
|
||||||
|
}
|
||||||
|
|
||||||
|
llama_kv_cache_msa_context::llama_kv_cache_msa_context(
|
||||||
|
llama_kv_cache_msa * kv,
|
||||||
|
slot_info_vec_t sinfos_base,
|
||||||
|
slot_info_vec_t sinfos_idx,
|
||||||
|
std::vector<llama_ubatch> ubatches) :
|
||||||
|
kv(kv),
|
||||||
|
ubatches(std::move(ubatches)),
|
||||||
|
// here we copy the ubatches. not sure if this is ideal
|
||||||
|
ctx_base(new llama_kv_cache_context(kv->get_base(), std::move(sinfos_base), this->ubatches)),
|
||||||
|
ctx_idx (new llama_kv_cache_context(kv->get_idx (), std::move(sinfos_idx), this->ubatches)),
|
||||||
|
status(llama_memory_status_combine(ctx_base->get_status(), ctx_idx->get_status())) {
|
||||||
|
}
|
||||||
|
|
||||||
|
llama_kv_cache_msa_context::~llama_kv_cache_msa_context() = default;
|
||||||
|
|
||||||
|
bool llama_kv_cache_msa_context::next() {
|
||||||
|
assert(status == LLAMA_MEMORY_STATUS_SUCCESS);
|
||||||
|
|
||||||
|
ctx_base->next();
|
||||||
|
ctx_idx ->next();
|
||||||
|
|
||||||
|
if (++i_next >= ubatches.size()) {
|
||||||
|
return false;
|
||||||
|
}
|
||||||
|
|
||||||
|
return true;
|
||||||
|
}
|
||||||
|
|
||||||
|
bool llama_kv_cache_msa_context::apply() {
|
||||||
|
assert(!llama_memory_status_is_fail(status));
|
||||||
|
|
||||||
|
bool res = true;
|
||||||
|
|
||||||
|
res = res & ctx_base->apply();
|
||||||
|
res = res & ctx_idx ->apply();
|
||||||
|
|
||||||
|
return res;
|
||||||
|
}
|
||||||
|
|
||||||
|
llama_memory_status llama_kv_cache_msa_context::get_status() const {
|
||||||
|
return status;
|
||||||
|
}
|
||||||
|
|
||||||
|
const llama_ubatch & llama_kv_cache_msa_context::get_ubatch() const {
|
||||||
|
assert(status == LLAMA_MEMORY_STATUS_SUCCESS);
|
||||||
|
|
||||||
|
return ubatches[i_next];
|
||||||
|
}
|
||||||
|
|
||||||
|
const llama_kv_cache_context * llama_kv_cache_msa_context::get_base() const {
|
||||||
|
assert(status == LLAMA_MEMORY_STATUS_SUCCESS);
|
||||||
|
|
||||||
|
return static_cast<const llama_kv_cache_context *>(ctx_base.get());
|
||||||
|
}
|
||||||
|
|
||||||
|
const llama_kv_cache_context * llama_kv_cache_msa_context::get_idx() const {
|
||||||
|
assert(status == LLAMA_MEMORY_STATUS_SUCCESS);
|
||||||
|
|
||||||
|
return static_cast<const llama_kv_cache_context *>(ctx_idx.get());
|
||||||
|
}
|
||||||
|
|
||||||
|
uint32_t llama_kv_cache_msa_context::get_n_pos() const {
|
||||||
|
// pad the value so that the graph remains constant across batches and can be reused
|
||||||
|
const uint32_t n_pad_cur = std::max(kv->get_n_pad(), 256u);
|
||||||
|
|
||||||
|
llama_pos pos_max = -1;
|
||||||
|
|
||||||
|
for (llama_seq_id seq_id = 0; seq_id < (llama_seq_id) kv->get_n_seq_max(); ++seq_id) {
|
||||||
|
pos_max = std::max(pos_max, kv->seq_pos_max(seq_id));
|
||||||
|
}
|
||||||
|
|
||||||
|
return std::max(n_pad_cur, GGML_PAD((uint32_t) (pos_max + 1), n_pad_cur));
|
||||||
|
}
|
||||||
|
|
||||||
|
void llama_kv_cache_msa_context::set_input_cell_pos(ggml_tensor * dst, const llama_ubatch * ubatch, int32_t div) const {
|
||||||
|
GGML_ASSERT(ggml_backend_buffer_is_host(dst->buffer));
|
||||||
|
GGML_ASSERT(dst->type == GGML_TYPE_I32);
|
||||||
|
GGML_ASSERT(div > 0);
|
||||||
|
|
||||||
|
const int64_t n_tokens = ubatch->n_tokens;
|
||||||
|
const int64_t n_kv = dst->ne[0];
|
||||||
|
const int64_t n_stream_ub = dst->ne[1];
|
||||||
|
|
||||||
|
GGML_ASSERT(n_tokens % n_stream_ub == 0);
|
||||||
|
const int64_t n_tps = n_tokens/n_stream_ub;
|
||||||
|
|
||||||
|
int32_t * data = (int32_t *) dst->data;
|
||||||
|
|
||||||
|
for (int64_t s = 0; s < n_stream_ub; ++s) {
|
||||||
|
const llama_seq_id seq_id = ubatch->seq_id[s*n_tps][0];
|
||||||
|
|
||||||
|
const auto & cells = kv->get_base()->get_cells(seq_id);
|
||||||
|
|
||||||
|
for (int64_t j = 0; j < n_kv; ++j) {
|
||||||
|
// the value for empty or other-sequence cells is irrelevant as consumers mask them
|
||||||
|
data[s*n_kv + j] =
|
||||||
|
cells.is_empty(j) || !cells.seq_has(j, seq_id)
|
||||||
|
? 0
|
||||||
|
: (int32_t) (cells.pos_get(j)/div);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
void llama_kv_cache_msa_context::set_input_pos_slot(ggml_tensor * dst, const llama_ubatch * ubatch) const {
|
||||||
|
GGML_ASSERT(ggml_backend_buffer_is_host(dst->buffer));
|
||||||
|
GGML_ASSERT(dst->type == GGML_TYPE_I32 || dst->type == GGML_TYPE_F32);
|
||||||
|
|
||||||
|
const int64_t n_tokens = ubatch->n_tokens;
|
||||||
|
const int64_t n_pos = dst->ne[0];
|
||||||
|
const int64_t n_stream_ub = dst->ne[1];
|
||||||
|
|
||||||
|
GGML_ASSERT(n_tokens % n_stream_ub == 0);
|
||||||
|
const int64_t n_tps = n_tokens/n_stream_ub;
|
||||||
|
|
||||||
|
for (int64_t s = 0; s < n_stream_ub; ++s) {
|
||||||
|
const llama_seq_id seq_id = ubatch->seq_id[s*n_tps][0];
|
||||||
|
|
||||||
|
const auto & cells = kv->get_base()->get_cells(seq_id);
|
||||||
|
|
||||||
|
std::vector<int32_t> map(n_pos, 0);
|
||||||
|
|
||||||
|
for (uint32_t j = 0; j < cells.size(); ++j) {
|
||||||
|
if (cells.is_empty(j) || !cells.seq_has(j, seq_id)) {
|
||||||
|
continue;
|
||||||
|
}
|
||||||
|
|
||||||
|
const llama_pos p0 = cells.pos_get(j);
|
||||||
|
|
||||||
|
if (p0 < 0 || p0 >= n_pos) {
|
||||||
|
continue;
|
||||||
|
}
|
||||||
|
|
||||||
|
map[p0] = (int32_t) j;
|
||||||
|
}
|
||||||
|
|
||||||
|
if (dst->type == GGML_TYPE_I32) {
|
||||||
|
int32_t * data = (int32_t *) dst->data + s*n_pos;
|
||||||
|
std::copy(map.begin(), map.end(), data);
|
||||||
|
} else {
|
||||||
|
float * data = (float *) dst->data + s*n_pos;
|
||||||
|
for (int64_t p = 0; p < n_pos; ++p) {
|
||||||
|
data[p] = (float) map[p];
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
void llama_kv_cache_msa_context::set_input_pos_mask(ggml_tensor * dst, const llama_ubatch * ubatch) const {
|
||||||
|
GGML_ASSERT(ggml_backend_buffer_is_host(dst->buffer));
|
||||||
|
GGML_ASSERT(dst->type == GGML_TYPE_F32);
|
||||||
|
|
||||||
|
const int64_t n_tokens = ubatch->n_tokens;
|
||||||
|
const int64_t n_pos = dst->ne[0];
|
||||||
|
|
||||||
|
GGML_ASSERT(dst->ne[1] == n_tokens);
|
||||||
|
|
||||||
|
const uint32_t n_swa = kv->get_n_swa();
|
||||||
|
const llama_swa_type swa_type = kv->get_swa_type();
|
||||||
|
|
||||||
|
float * data = (float *) dst->data;
|
||||||
|
|
||||||
|
std::fill(data, data + n_pos*n_tokens, -INFINITY);
|
||||||
|
|
||||||
|
for (int64_t i = 0; i < n_tokens; ++i) {
|
||||||
|
const llama_seq_id seq_id = ubatch->seq_id[i][0];
|
||||||
|
|
||||||
|
const auto & cells = kv->get_base()->get_cells(seq_id);
|
||||||
|
|
||||||
|
const llama_pos p1 = ubatch->pos[i];
|
||||||
|
|
||||||
|
for (uint32_t j = 0; j < cells.size(); ++j) {
|
||||||
|
if (cells.is_empty(j) || !cells.seq_has(j, seq_id)) {
|
||||||
|
continue;
|
||||||
|
}
|
||||||
|
|
||||||
|
const llama_pos p0 = cells.pos_get(j);
|
||||||
|
|
||||||
|
if (p0 < 0 || p0 >= n_pos) {
|
||||||
|
continue;
|
||||||
|
}
|
||||||
|
|
||||||
|
// causal mask
|
||||||
|
if (p0 > p1) {
|
||||||
|
continue;
|
||||||
|
}
|
||||||
|
|
||||||
|
// apply SWA if any
|
||||||
|
if (llama_hparams::is_masked_swa(n_swa, swa_type, p0, p1)) {
|
||||||
|
continue;
|
||||||
|
}
|
||||||
|
|
||||||
|
data[i*n_pos + p0] = 0.0f;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,153 @@
|
|||||||
|
#pragma once
|
||||||
|
|
||||||
|
#include "llama-kv-cache.h"
|
||||||
|
|
||||||
|
#include <vector>
|
||||||
|
|
||||||
|
// llama_kv_cache_msa
|
||||||
|
|
||||||
|
// uses two instances of llama_kv_cache, one for K/V tensors, and one for the MSA indexer tensors
|
||||||
|
// both receive identical sequence operations and identical ubatches, so their cell layouts stay in synced.
|
||||||
|
// the context also exposes per-ubatch pos - cell translation maps populated from llama_kv_cells via
|
||||||
|
// llama_kv_cache::get_cells(), which the model graph uses to run MSA block selection in position space
|
||||||
|
|
||||||
|
class llama_kv_cache_msa : public llama_memory_i {
|
||||||
|
public:
|
||||||
|
llama_kv_cache_msa(
|
||||||
|
const llama_model & model,
|
||||||
|
ggml_type type_k,
|
||||||
|
ggml_type type_v,
|
||||||
|
bool v_trans,
|
||||||
|
bool offload,
|
||||||
|
bool unified,
|
||||||
|
uint32_t kv_size,
|
||||||
|
uint32_t n_seq_max,
|
||||||
|
uint32_t n_pad,
|
||||||
|
uint32_t n_swa,
|
||||||
|
llama_swa_type swa_type,
|
||||||
|
const layer_filter_cb & filter,
|
||||||
|
const layer_filter_cb & filter_idx,
|
||||||
|
const layer_reuse_cb & reuse);
|
||||||
|
|
||||||
|
~llama_kv_cache_msa() = default;
|
||||||
|
|
||||||
|
// llama_memory_i
|
||||||
|
|
||||||
|
llama_memory_context_ptr init_batch(
|
||||||
|
llama_batch_allocr & balloc,
|
||||||
|
uint32_t n_ubatch,
|
||||||
|
bool embd_all) override;
|
||||||
|
|
||||||
|
llama_memory_context_ptr init_full() override;
|
||||||
|
|
||||||
|
llama_memory_context_ptr init_update(llama_context * lctx, bool optimize) override;
|
||||||
|
|
||||||
|
bool get_can_shift() const override;
|
||||||
|
|
||||||
|
void clear(bool data) override;
|
||||||
|
|
||||||
|
bool seq_rm (llama_seq_id seq_id, llama_pos p0, llama_pos p1) override;
|
||||||
|
void seq_cp (llama_seq_id seq_id_src, llama_seq_id seq_id_dst, llama_pos p0, llama_pos p1) override;
|
||||||
|
void seq_keep(llama_seq_id seq_id) override;
|
||||||
|
void seq_add (llama_seq_id seq_id, llama_pos p0, llama_pos p1, llama_pos shift) override;
|
||||||
|
void seq_div (llama_seq_id seq_id, llama_pos p0, llama_pos p1, int d) override;
|
||||||
|
|
||||||
|
llama_pos seq_pos_min(llama_seq_id seq_id) const override;
|
||||||
|
llama_pos seq_pos_max(llama_seq_id seq_id) const override;
|
||||||
|
|
||||||
|
std::map<ggml_backend_buffer_type_t, size_t> memory_breakdown() const override;
|
||||||
|
|
||||||
|
// state write/load
|
||||||
|
|
||||||
|
void state_write(llama_io_write_i & io, llama_seq_id seq_id = -1, llama_state_seq_flags flags = 0) const override;
|
||||||
|
void state_read (llama_io_read_i & io, llama_seq_id seq_id = -1, llama_state_seq_flags flags = 0) override;
|
||||||
|
|
||||||
|
// llama_kv_cache_msa specific API
|
||||||
|
|
||||||
|
llama_kv_cache * get_base() const;
|
||||||
|
llama_kv_cache * get_idx () const;
|
||||||
|
|
||||||
|
uint32_t get_n_pad() const { return n_pad; }
|
||||||
|
uint32_t get_n_seq_max() const { return n_seq_max; }
|
||||||
|
uint32_t get_n_swa() const { return n_swa; }
|
||||||
|
llama_swa_type get_swa_type() const { return swa_type; }
|
||||||
|
|
||||||
|
private:
|
||||||
|
// keep the indexer KV cache hparams instance here as llama_kv_cache stores only a reference
|
||||||
|
llama_hparams hparams_idx;
|
||||||
|
|
||||||
|
const uint32_t n_stream = 1;
|
||||||
|
const uint32_t n_seq_max = 1;
|
||||||
|
const uint32_t n_pad = 1;
|
||||||
|
|
||||||
|
const uint32_t n_swa = 0;
|
||||||
|
const llama_swa_type swa_type = LLAMA_SWA_TYPE_NONE;
|
||||||
|
|
||||||
|
std::unique_ptr<llama_kv_cache> kv_base;
|
||||||
|
std::unique_ptr<llama_kv_cache> kv_idx;
|
||||||
|
};
|
||||||
|
|
||||||
|
class llama_kv_cache_msa_context : public llama_memory_context_i {
|
||||||
|
public:
|
||||||
|
using slot_info_vec_t = llama_kv_cache::slot_info_vec_t;
|
||||||
|
|
||||||
|
// used for errors
|
||||||
|
llama_kv_cache_msa_context(llama_memory_status status);
|
||||||
|
|
||||||
|
// used to create a full-cache context
|
||||||
|
llama_kv_cache_msa_context(
|
||||||
|
llama_kv_cache_msa * kv);
|
||||||
|
|
||||||
|
// used to create an update context
|
||||||
|
llama_kv_cache_msa_context(
|
||||||
|
llama_kv_cache_msa * kv,
|
||||||
|
llama_context * lctx,
|
||||||
|
bool optimize);
|
||||||
|
|
||||||
|
// used to create a batch processing context from a batch
|
||||||
|
llama_kv_cache_msa_context(
|
||||||
|
llama_kv_cache_msa * kv,
|
||||||
|
slot_info_vec_t sinfos_base,
|
||||||
|
slot_info_vec_t sinfos_idx,
|
||||||
|
std::vector<llama_ubatch> ubatches);
|
||||||
|
|
||||||
|
virtual ~llama_kv_cache_msa_context();
|
||||||
|
|
||||||
|
// llama_memory_context_i
|
||||||
|
|
||||||
|
bool next() override;
|
||||||
|
bool apply() override;
|
||||||
|
|
||||||
|
llama_memory_status get_status() const override;
|
||||||
|
const llama_ubatch & get_ubatch() const override;
|
||||||
|
|
||||||
|
// llama_kv_cache_msa_context specific API
|
||||||
|
|
||||||
|
const llama_kv_cache_context * get_base() const;
|
||||||
|
const llama_kv_cache_context * get_idx () const;
|
||||||
|
|
||||||
|
// max position currently present in the cache plus one, padded MSA blocks are defined over token positions
|
||||||
|
// so the block-selection tensors are sized by this value rather than by the number of cells
|
||||||
|
uint32_t get_n_pos() const;
|
||||||
|
|
||||||
|
// position <-> cell translation maps, populated from the base cache cells
|
||||||
|
// the model graph relates cache contents to token positions only through these per ubatch inputs
|
||||||
|
// value for empty or other-sequence cells is 0 so consumers must mask them
|
||||||
|
void set_input_cell_pos(ggml_tensor * dst, const llama_ubatch * ubatch, int32_t div) const;
|
||||||
|
// positions without a cell map to cell 0, consumers must mask them assumes one sequence per stream
|
||||||
|
void set_input_pos_slot(ggml_tensor * dst, const llama_ubatch * ubatch) const;
|
||||||
|
void set_input_pos_mask(ggml_tensor * dst, const llama_ubatch * ubatch) const;
|
||||||
|
|
||||||
|
private:
|
||||||
|
llama_kv_cache_msa * kv;
|
||||||
|
|
||||||
|
// the index of the next ubatch to process
|
||||||
|
size_t i_next = 0;
|
||||||
|
|
||||||
|
std::vector<llama_ubatch> ubatches;
|
||||||
|
|
||||||
|
const llama_memory_context_ptr ctx_base;
|
||||||
|
const llama_memory_context_ptr ctx_idx;
|
||||||
|
|
||||||
|
const llama_memory_status status;
|
||||||
|
};
|
||||||
+20
-278
@@ -112,7 +112,7 @@ llama_kv_cache::llama_kv_cache(
|
|||||||
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(3u*(1 + n_stream)*n_layer*ggml_tensor_overhead()), //Reserve tensor metadata for up to 3 tensors per layer (K, V, and optional K_idx), plus one view per tensor per stream.
|
/*.mem_size =*/ size_t(2u*(1 + n_stream)*n_layer*ggml_tensor_overhead()),
|
||||||
/*.mem_buffer =*/ NULL,
|
/*.mem_buffer =*/ NULL,
|
||||||
/*.no_alloc =*/ true,
|
/*.no_alloc =*/ true,
|
||||||
};
|
};
|
||||||
@@ -242,25 +242,9 @@ llama_kv_cache::llama_kv_cache(
|
|||||||
v_stream.push_back(has_v ? ggml_view_2d(ctx, v, n_embd_v_gqa, kv_size, v->nb[1], s*v->nb[2]) : nullptr);
|
v_stream.push_back(has_v ? ggml_view_2d(ctx, v, n_embd_v_gqa, kv_size, v->nb[1], s*v->nb[2]) : nullptr);
|
||||||
}
|
}
|
||||||
|
|
||||||
const uint32_t n_embd_k_idx = hparams.n_embd_k_idx(il);
|
|
||||||
ggml_tensor * k_idx = n_embd_k_idx > 0
|
|
||||||
? ggml_new_tensor_3d(ctx, GGML_TYPE_F32, n_embd_k_idx, kv_size, n_stream)
|
|
||||||
: nullptr;
|
|
||||||
if (k_idx) {
|
|
||||||
ggml_format_name(k_idx, "cache_k_idx_l%d", il);
|
|
||||||
msa_strict_slots = (n_stream == n_seq_max);
|
|
||||||
}
|
|
||||||
|
|
||||||
std::vector<ggml_tensor *> k_idx_stream;
|
|
||||||
for (uint32_t s = 0; s < n_stream; ++s) {
|
|
||||||
k_idx_stream.push_back(k_idx
|
|
||||||
? ggml_view_2d(ctx, k_idx, n_embd_k_idx, kv_size, k_idx->nb[1], s*k_idx->nb[2])
|
|
||||||
: nullptr);
|
|
||||||
}
|
|
||||||
|
|
||||||
map_layer_ids[il] = layers.size();
|
map_layer_ids[il] = layers.size();
|
||||||
|
|
||||||
layers.push_back({ il, k, v, k_idx, k_stream, v_stream, k_idx_stream });
|
layers.push_back({ il, k, v, k_stream, v_stream, });
|
||||||
}
|
}
|
||||||
|
|
||||||
if (reuse) {
|
if (reuse) {
|
||||||
@@ -309,24 +293,13 @@ llama_kv_cache::llama_kv_cache(
|
|||||||
}
|
}
|
||||||
|
|
||||||
{
|
{
|
||||||
const size_t memory_size_k = size_k_bytes();
|
const size_t memory_size_k = size_k_bytes();
|
||||||
const size_t memory_size_v = size_v_bytes();
|
const size_t memory_size_v = size_v_bytes();
|
||||||
const size_t memory_size_k_idx = size_k_idx_bytes();
|
|
||||||
const size_t memory_size_total = memory_size_k + memory_size_v + memory_size_k_idx;
|
|
||||||
|
|
||||||
constexpr float mib = 1024.0f * 1024.0f;
|
LLAMA_LOG_INFO("%s: size = %7.2f MiB (%6u cells, %3d layers, %2u/%u seqs), K (%s): %7.2f MiB, V (%s): %7.2f MiB\n", __func__,
|
||||||
|
(float)(memory_size_k + memory_size_v) / (1024.0f * 1024.0f), kv_size, (int) layers.size(), n_seq_max, n_stream,
|
||||||
const std::string k_log = format(", K (%s): %7.2f MiB", ggml_type_name(type_k), (float) memory_size_k / mib);
|
ggml_type_name(type_k), (float)memory_size_k / (1024.0f * 1024.0f),
|
||||||
const std::string v_log = format(", V (%s): %7.2f MiB", ggml_type_name(type_v), (float) memory_size_v / mib);
|
ggml_type_name(type_v), (float)memory_size_v / (1024.0f * 1024.0f));
|
||||||
|
|
||||||
std::string k_idx_log;
|
|
||||||
if (memory_size_k_idx > 0) {
|
|
||||||
k_idx_log = format(", K_idx (%s): %7.2f MiB", ggml_type_name(GGML_TYPE_F32), (float) memory_size_k_idx / mib);
|
|
||||||
}
|
|
||||||
|
|
||||||
LLAMA_LOG_INFO("%s: size = %7.2f MiB (%6u cells, %3d layers, %2u/%u seqs)%s%s%s\n", __func__,
|
|
||||||
(float) memory_size_total / mib, kv_size, (int) layers.size(), n_seq_max, n_stream,
|
|
||||||
k_log.c_str(), v_log.c_str(), k_idx_log.c_str());
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// TODO: refactor [TAG_KV_CACHE_SHARE_CELLS]
|
// TODO: refactor [TAG_KV_CACHE_SHARE_CELLS]
|
||||||
@@ -419,39 +392,6 @@ bool llama_kv_cache::seq_rm(llama_seq_id seq_id, llama_pos p0, llama_pos p1) {
|
|||||||
p1 = std::numeric_limits<llama_pos>::max();
|
p1 = std::numeric_limits<llama_pos>::max();
|
||||||
}
|
}
|
||||||
|
|
||||||
// empty range - nothing to remove
|
|
||||||
if (p0 >= p1) {
|
|
||||||
return true;
|
|
||||||
}
|
|
||||||
|
|
||||||
// MSA anchors block selection to absolute cache slots (slot == position). Tail trim and full removal preserve this invariant, but removing a prefix
|
|
||||||
// or middle range would free slots while later cells survive, desynchronizing the indexer cache. Reject such removals before modifying the cache.
|
|
||||||
if (msa_strict_slots) {
|
|
||||||
for (llama_seq_id sid = 0; sid < (llama_seq_id) seq_to_stream.size(); ++sid) {
|
|
||||||
if (seq_id >= 0 && sid != seq_id) {
|
|
||||||
continue;
|
|
||||||
}
|
|
||||||
|
|
||||||
const auto & cells = v_cells[seq_to_stream[sid]];
|
|
||||||
|
|
||||||
const llama_pos pmin = cells.seq_pos_min(sid);
|
|
||||||
const llama_pos pmax = cells.seq_pos_max(sid);
|
|
||||||
|
|
||||||
if (pmin < 0) {
|
|
||||||
continue; // empty sequence
|
|
||||||
}
|
|
||||||
|
|
||||||
const bool overlaps = p0 <= pmax && p1 > pmin; // the range removes something
|
|
||||||
const bool leaves_tail = p1 <= pmax; // cells beyond the range survive
|
|
||||||
|
|
||||||
if (overlaps && leaves_tail) {
|
|
||||||
LLAMA_LOG_WARN("%s: MSA: partial (non-suffix) removal [%d, %d) for seq %d is not supported "
|
|
||||||
"(block selection is anchored to cache slots) - rejected\n", __func__, p0, p1, sid);
|
|
||||||
return false;
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
if (seq_id >= 0) {
|
if (seq_id >= 0) {
|
||||||
auto & cells = v_cells[seq_to_stream[seq_id]];
|
auto & cells = v_cells[seq_to_stream[seq_id]];
|
||||||
auto & head = v_heads[seq_to_stream[seq_id]];
|
auto & head = v_heads[seq_to_stream[seq_id]];
|
||||||
@@ -906,10 +846,6 @@ bool llama_kv_cache::update(llama_context * lctx, bool do_shift, const stream_co
|
|||||||
if (layer.v_stream[ssrc]) {
|
if (layer.v_stream[ssrc]) {
|
||||||
ggml_backend_tensor_copy(layer.v_stream[ssrc], layer.v_stream[sdst]);
|
ggml_backend_tensor_copy(layer.v_stream[ssrc], layer.v_stream[sdst]);
|
||||||
}
|
}
|
||||||
if (layer.k_idx_stream[ssrc]) {
|
|
||||||
GGML_ASSERT(layer.k_idx_stream[sdst]);
|
|
||||||
ggml_backend_tensor_copy(layer.k_idx_stream[ssrc], layer.k_idx_stream[sdst]);
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -1058,44 +994,6 @@ llama_kv_cache::slot_info llama_kv_cache::find_slot(const llama_ubatch & ubatch,
|
|||||||
|
|
||||||
const auto & cells = v_cells[seq_to_stream[seq_id]];
|
const auto & cells = v_cells[seq_to_stream[seq_id]];
|
||||||
|
|
||||||
if (n_tokens > cells.size()) {
|
|
||||||
LLAMA_LOG_ERROR("%s: n_tokens = %d > size = %u\n", __func__, n_tokens, cells.size());
|
|
||||||
return { };
|
|
||||||
}
|
|
||||||
|
|
||||||
// MSA block selection assumes slot == logical position (append-only streams).
|
|
||||||
if (msa_strict_slots) {
|
|
||||||
for (uint32_t ii = 0; ii < n_tokens; ++ii) {
|
|
||||||
const llama_pos pos = ubatch.pos[s*n_tokens + ii];
|
|
||||||
|
|
||||||
if (pos < 0 || (uint64_t) pos >= cells.size()) {
|
|
||||||
LLAMA_LOG_WARN("%s: MSA: position %d is outside the cache range [0, %u)\n",
|
|
||||||
__func__, pos, cells.size());
|
|
||||||
return { };
|
|
||||||
}
|
|
||||||
|
|
||||||
const uint32_t idx = (uint32_t) pos;
|
|
||||||
|
|
||||||
if (!cells.is_empty(idx)) {
|
|
||||||
LLAMA_LOG_WARN("%s: MSA: required slot %u is already occupied (stream %u)\n",
|
|
||||||
__func__, idx, seq_to_stream[seq_id]);
|
|
||||||
return { };
|
|
||||||
}
|
|
||||||
|
|
||||||
// strictly increasing positions, rules out duplicates and, for contiguous requests, is tightened to exact adjacency
|
|
||||||
if (!res.idxs[s].empty() && (cont ? idx != res.idxs[s].back() + 1
|
|
||||||
: idx <= res.idxs[s].back())) {
|
|
||||||
LLAMA_LOG_WARN("%s: MSA: token positions are not %s within the ubatch\n",
|
|
||||||
__func__, cont ? "contiguous" : "strictly increasing");
|
|
||||||
return { };
|
|
||||||
}
|
|
||||||
|
|
||||||
res.idxs[s].push_back(idx);
|
|
||||||
}
|
|
||||||
|
|
||||||
continue;
|
|
||||||
}
|
|
||||||
|
|
||||||
uint32_t head_cur = v_heads[seq_to_stream[seq_id]];
|
uint32_t head_cur = v_heads[seq_to_stream[seq_id]];
|
||||||
|
|
||||||
// if we have enough unused cells before the current head ->
|
// if we have enough unused cells before the current head ->
|
||||||
@@ -1104,6 +1002,11 @@ llama_kv_cache::slot_info llama_kv_cache::find_slot(const llama_ubatch & ubatch,
|
|||||||
head_cur = 0;
|
head_cur = 0;
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if (n_tokens > cells.size()) {
|
||||||
|
LLAMA_LOG_ERROR("%s: n_tokens = %d > size = %u\n", __func__, n_tokens, cells.size());
|
||||||
|
return { };
|
||||||
|
}
|
||||||
|
|
||||||
uint32_t n_tested = 0;
|
uint32_t n_tested = 0;
|
||||||
|
|
||||||
// for continuous slots, we test that all tokens in the ubatch fit, starting from the current head
|
// for continuous slots, we test that all tokens in the ubatch fit, starting from the current head
|
||||||
@@ -1210,15 +1113,6 @@ void llama_kv_cache::apply_ubatch(const slot_info & sinfo, const llama_ubatch &
|
|||||||
|
|
||||||
const auto idx = sinfo.idxs[s][ii];
|
const auto idx = sinfo.idxs[s][ii];
|
||||||
|
|
||||||
if (msa_strict_slots && (llama_pos) idx != ubatch.pos[i]) {
|
|
||||||
LLAMA_LOG_ERROR("%s: MSA slot/position invariant violated: "
|
|
||||||
"writing pos %d into cell %u (stream %u). The indexer cache "
|
|
||||||
"would desync and block selection would silently corrupt. "
|
|
||||||
"This is a bug, please report it with reproduction steps.\n",
|
|
||||||
__func__, ubatch.pos[i], idx, sinfo.strm[s]);
|
|
||||||
GGML_ABORT("MSA: slot != pos");
|
|
||||||
}
|
|
||||||
|
|
||||||
if (!cells.is_empty(idx)) {
|
if (!cells.is_empty(idx)) {
|
||||||
assert(cells.seq_count(idx) == 1);
|
assert(cells.seq_count(idx) == 1);
|
||||||
|
|
||||||
@@ -1262,8 +1156,7 @@ void llama_kv_cache::apply_ubatch(const slot_info & sinfo, const llama_ubatch &
|
|||||||
LLAMA_LOG_DEBUG("%s: purging positions [%d, %d] of sequence %d from KV cache\n",
|
LLAMA_LOG_DEBUG("%s: purging positions [%d, %d] of sequence %d from KV cache\n",
|
||||||
__func__, cells.seq_pos_min(s), seq_pos_max_rm[s], s);
|
__func__, cells.seq_pos_min(s), seq_pos_max_rm[s], s);
|
||||||
|
|
||||||
// under MSA strict slots this path should be unreachable, since strict MSA placement never selects occupied cells
|
seq_rm(s, cells.seq_pos_min(s), seq_pos_max_rm[s] + 1);
|
||||||
GGML_ASSERT(seq_rm(s, cells.seq_pos_min(s), seq_pos_max_rm[s] + 1));
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -1283,12 +1176,6 @@ bool llama_kv_cache::get_can_shift() const {
|
|||||||
if (hparams.n_pos_per_embd() > 1) {
|
if (hparams.n_pos_per_embd() > 1) {
|
||||||
return false;
|
return false;
|
||||||
}
|
}
|
||||||
// shifting would leave k_idx stale
|
|
||||||
for (const auto & layer : layers) {
|
|
||||||
if (layer.k_idx) {
|
|
||||||
return false;
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return true;
|
return true;
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -1337,6 +1224,12 @@ ggml_tensor * llama_kv_cache::get_k_storage(int32_t il) const {
|
|||||||
return layers[ikv].k;
|
return layers[ikv].k;
|
||||||
}
|
}
|
||||||
|
|
||||||
|
const llama_kv_cells & llama_kv_cache::get_cells(llama_seq_id seq_id) const {
|
||||||
|
GGML_ASSERT(seq_id >= 0 && (size_t) seq_id < seq_to_stream.size());
|
||||||
|
|
||||||
|
return v_cells[seq_to_stream[seq_id]];
|
||||||
|
}
|
||||||
|
|
||||||
uint32_t llama_kv_cache::get_n_kv(const slot_info & sinfo) const {
|
uint32_t llama_kv_cache::get_n_kv(const slot_info & sinfo) const {
|
||||||
uint32_t result = 0;
|
uint32_t result = 0;
|
||||||
|
|
||||||
@@ -1405,23 +1298,6 @@ ggml_tensor * llama_kv_cache::get_v(ggml_context * ctx, int32_t il, uint32_t n_k
|
|||||||
ggml_row_size(v->type, kv_size*n_embd_v_gqa)*sinfo.s0);
|
ggml_row_size(v->type, kv_size*n_embd_v_gqa)*sinfo.s0);
|
||||||
}
|
}
|
||||||
|
|
||||||
ggml_tensor * llama_kv_cache::get_k_idx(ggml_context * ctx, int32_t il, uint32_t n_kv, const slot_info & sinfo) const {
|
|
||||||
const int32_t ikv = map_layer_ids.at(il);
|
|
||||||
auto * k_idx = layers[ikv].k_idx;
|
|
||||||
GGML_ASSERT(k_idx);
|
|
||||||
|
|
||||||
const uint64_t kv_size = get_size();
|
|
||||||
const int64_t n_idx = k_idx->ne[0]; // 128
|
|
||||||
const uint32_t ns = sinfo.s1 - sinfo.s0 + 1;
|
|
||||||
|
|
||||||
return ggml_view_4d(ctx, k_idx,
|
|
||||||
n_idx, 1, n_kv, ns,
|
|
||||||
ggml_row_size(k_idx->type, n_idx), // nb1 (single head)
|
|
||||||
ggml_row_size(k_idx->type, n_idx), // nb2 (per cell)
|
|
||||||
ggml_row_size(k_idx->type, n_idx*kv_size), // nb3 (per stream)
|
|
||||||
ggml_row_size(k_idx->type, n_idx*kv_size)*sinfo.s0);
|
|
||||||
}
|
|
||||||
|
|
||||||
ggml_tensor * llama_kv_cache::cpy_k(ggml_context * ctx, ggml_tensor * k_cur, ggml_tensor * k_idxs, int32_t il, const slot_info & sinfo) const {
|
ggml_tensor * llama_kv_cache::cpy_k(ggml_context * ctx, ggml_tensor * k_cur, ggml_tensor * k_idxs, int32_t il, const slot_info & sinfo) const {
|
||||||
GGML_UNUSED(sinfo);
|
GGML_UNUSED(sinfo);
|
||||||
|
|
||||||
@@ -1523,28 +1399,6 @@ ggml_tensor * llama_kv_cache::build_input_k_idxs(ggml_context * ctx, const llama
|
|||||||
return k_idxs;
|
return k_idxs;
|
||||||
}
|
}
|
||||||
|
|
||||||
ggml_tensor * llama_kv_cache::cpy_k_idx(ggml_context * ctx, ggml_tensor * k_idx_cur, ggml_tensor * k_idxs, int32_t il, const slot_info & sinfo) const {
|
|
||||||
GGML_UNUSED(sinfo);
|
|
||||||
const int32_t ikv = map_layer_ids.at(il);
|
|
||||||
ggml_tensor * k_idx = layers[ikv].k_idx;
|
|
||||||
GGML_ASSERT(k_idx && "cpy_k_idx on a layer with no indexer cache");
|
|
||||||
|
|
||||||
const int64_t n_embd_head = k_idx_cur->ne[0]; // 128
|
|
||||||
const int64_t n_head = k_idx_cur->ne[1]; // 1
|
|
||||||
const int64_t n_tokens = k_idx_cur->ne[2];
|
|
||||||
const int64_t n_embd_gqa = n_embd_head*n_head; // 128
|
|
||||||
|
|
||||||
GGML_ASSERT(ggml_row_size(k_idx_cur->type, n_embd_head) == k_idx_cur->nb[1]);
|
|
||||||
k_idx_cur = ggml_view_2d(ctx, k_idx_cur, n_embd_gqa, n_tokens, k_idx_cur->nb[2], 0);
|
|
||||||
|
|
||||||
const int64_t n_stream = k_idx->ne[2];
|
|
||||||
if (n_stream > 1) {
|
|
||||||
const int64_t kv_size = get_size();
|
|
||||||
k_idx = ggml_reshape_2d(ctx, k_idx, n_embd_gqa, kv_size*n_stream);
|
|
||||||
}
|
|
||||||
return ggml_set_rows(ctx, k_idx, k_idx_cur, k_idxs); // same k_idxs as the K store
|
|
||||||
}
|
|
||||||
|
|
||||||
ggml_tensor * llama_kv_cache::build_input_v_idxs(ggml_context * ctx, const llama_ubatch & ubatch) const {
|
ggml_tensor * llama_kv_cache::build_input_v_idxs(ggml_context * ctx, const llama_ubatch & ubatch) const {
|
||||||
const uint32_t n_tokens = ubatch.n_tokens;
|
const uint32_t n_tokens = ubatch.n_tokens;
|
||||||
|
|
||||||
@@ -1979,18 +1833,6 @@ size_t llama_kv_cache::size_v_bytes() const {
|
|||||||
return size_v_bytes;
|
return size_v_bytes;
|
||||||
}
|
}
|
||||||
|
|
||||||
size_t llama_kv_cache::size_k_idx_bytes() const {
|
|
||||||
size_t size_k_idx_bytes = 0;
|
|
||||||
|
|
||||||
for (const auto & layer : layers) {
|
|
||||||
if (layer.k_idx) {
|
|
||||||
size_k_idx_bytes += ggml_nbytes(layer.k_idx);
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
return size_k_idx_bytes;
|
|
||||||
}
|
|
||||||
|
|
||||||
ggml_tensor * llama_kv_cache::build_rope_shift(
|
ggml_tensor * llama_kv_cache::build_rope_shift(
|
||||||
const llama_cparams & cparams,
|
const llama_cparams & cparams,
|
||||||
ggml_context * ctx,
|
ggml_context * ctx,
|
||||||
@@ -2303,36 +2145,6 @@ void llama_kv_cache::state_write_data(llama_io_write_i & io, const cell_ranges_t
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
if (size_k_idx_bytes() > 0) {
|
|
||||||
const uint32_t has_k_idx_u32 = 1;
|
|
||||||
io.write(&has_k_idx_u32, sizeof(has_k_idx_u32));
|
|
||||||
|
|
||||||
for (const auto & layer : layers) {
|
|
||||||
const uint32_t layer_has_k_idx = layer.k_idx ? 1 : 0;
|
|
||||||
io.write(&layer_has_k_idx, sizeof(layer_has_k_idx));
|
|
||||||
|
|
||||||
if (!layer_has_k_idx) {
|
|
||||||
continue;
|
|
||||||
}
|
|
||||||
|
|
||||||
GGML_ASSERT(layer.k_idx_stream[cr.strm]);
|
|
||||||
|
|
||||||
const int32_t k_idx_type_i = (int32_t) layer.k_idx->type;
|
|
||||||
io.write(&k_idx_type_i, sizeof(k_idx_type_i));
|
|
||||||
|
|
||||||
const uint64_t k_idx_size_row = ggml_row_size(layer.k_idx->type, layer.k_idx->ne[0]);
|
|
||||||
io.write(&k_idx_size_row, sizeof(k_idx_size_row));
|
|
||||||
|
|
||||||
for (const auto & range : cr.data) {
|
|
||||||
const size_t range_size = range.second - range.first;
|
|
||||||
const size_t buf_size = range_size * k_idx_size_row;
|
|
||||||
const size_t offset = range.first * k_idx_size_row;
|
|
||||||
|
|
||||||
io.write_tensor(layer.k_idx_stream[cr.strm], offset, buf_size);
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
if (!v_trans) {
|
if (!v_trans) {
|
||||||
for (const auto & layer : layers) {
|
for (const auto & layer : layers) {
|
||||||
const uint32_t il = layer.il;
|
const uint32_t il = layer.il;
|
||||||
@@ -2581,68 +2393,6 @@ bool llama_kv_cache::state_read_data(llama_io_read_i & io, uint32_t strm, uint32
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
if (size_k_idx_bytes() > 0) {
|
|
||||||
uint32_t has_k_idx_u32 = 0;
|
|
||||||
io.read(&has_k_idx_u32, sizeof(has_k_idx_u32));
|
|
||||||
|
|
||||||
if (has_k_idx_u32 != 1) {
|
|
||||||
LLAMA_LOG_ERROR("%s: missing k_idx data in KV cache state\n", __func__);
|
|
||||||
return false;
|
|
||||||
}
|
|
||||||
|
|
||||||
for (const auto & layer : layers) {
|
|
||||||
uint32_t layer_has_k_idx = 0;
|
|
||||||
io.read(&layer_has_k_idx, sizeof(layer_has_k_idx));
|
|
||||||
|
|
||||||
const uint32_t expected_layer_has_k_idx = layer.k_idx ? 1 : 0;
|
|
||||||
|
|
||||||
if (layer_has_k_idx != expected_layer_has_k_idx) {
|
|
||||||
LLAMA_LOG_ERROR(
|
|
||||||
"%s: mismatched k_idx state for layer: got %u, expected %u\n",
|
|
||||||
__func__, layer_has_k_idx, expected_layer_has_k_idx);
|
|
||||||
return false;
|
|
||||||
}
|
|
||||||
|
|
||||||
if (!layer_has_k_idx) {
|
|
||||||
continue;
|
|
||||||
}
|
|
||||||
|
|
||||||
GGML_ASSERT(layer.k_idx_stream[strm]);
|
|
||||||
|
|
||||||
int32_t k_idx_type_i = -1;
|
|
||||||
io.read(&k_idx_type_i, sizeof(k_idx_type_i));
|
|
||||||
|
|
||||||
if (k_idx_type_i != (int32_t) layer.k_idx->type) {
|
|
||||||
LLAMA_LOG_ERROR(
|
|
||||||
"%s: mismatched k_idx type: got %d, expected %d\n",
|
|
||||||
__func__, k_idx_type_i, (int32_t) layer.k_idx->type);
|
|
||||||
return false;
|
|
||||||
}
|
|
||||||
|
|
||||||
uint64_t k_idx_size_row = 0;
|
|
||||||
io.read(&k_idx_size_row, sizeof(k_idx_size_row));
|
|
||||||
|
|
||||||
const uint64_t expected_k_idx_size_row = ggml_row_size(layer.k_idx->type, layer.k_idx->ne[0]);
|
|
||||||
|
|
||||||
if (k_idx_size_row != expected_k_idx_size_row) {
|
|
||||||
LLAMA_LOG_ERROR(
|
|
||||||
"%s: mismatched k_idx row size: got %zu, expected %zu\n",
|
|
||||||
__func__, (size_t) k_idx_size_row, (size_t) expected_k_idx_size_row);
|
|
||||||
return false;
|
|
||||||
}
|
|
||||||
|
|
||||||
if (cell_count) {
|
|
||||||
if (sinfo.is_contiguous()) {
|
|
||||||
io.read_tensor(layer.k_idx_stream[strm], sinfo.head() * k_idx_size_row, cell_count * k_idx_size_row);
|
|
||||||
} else {
|
|
||||||
for (uint32_t i = 0; i < cell_count; ++i) {
|
|
||||||
io.read_tensor(layer.k_idx_stream[strm], sinfo.idxs[0][i] * k_idx_size_row, k_idx_size_row);
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
if (!this->v_trans) {
|
if (!this->v_trans) {
|
||||||
for (const auto & layer : layers) {
|
for (const auto & layer : layers) {
|
||||||
const uint32_t il = layer.il;
|
const uint32_t il = layer.il;
|
||||||
@@ -2844,10 +2594,6 @@ ggml_tensor * llama_kv_cache_context::get_v(ggml_context * ctx, int32_t il) cons
|
|||||||
return kv->get_v(ctx, il, n_kv, sinfos[i_cur]);
|
return kv->get_v(ctx, il, n_kv, sinfos[i_cur]);
|
||||||
}
|
}
|
||||||
|
|
||||||
ggml_tensor * llama_kv_cache_context::get_k_idx(ggml_context * ctx, int32_t il) const {
|
|
||||||
return kv->get_k_idx(ctx, il, n_kv, sinfos[i_cur]);
|
|
||||||
}
|
|
||||||
|
|
||||||
ggml_tensor * llama_kv_cache_context::cpy_k(ggml_context * ctx, ggml_tensor * k_cur, ggml_tensor * k_idxs, int32_t il) const {
|
ggml_tensor * llama_kv_cache_context::cpy_k(ggml_context * ctx, ggml_tensor * k_cur, ggml_tensor * k_idxs, int32_t il) const {
|
||||||
return kv->cpy_k(ctx, k_cur, k_idxs, il, sinfos[i_cur]);
|
return kv->cpy_k(ctx, k_cur, k_idxs, il, sinfos[i_cur]);
|
||||||
}
|
}
|
||||||
@@ -2856,10 +2602,6 @@ ggml_tensor * llama_kv_cache_context::cpy_v(ggml_context * ctx, ggml_tensor * v_
|
|||||||
return kv->cpy_v(ctx, v_cur, v_idxs, il, sinfos[i_cur]);
|
return kv->cpy_v(ctx, v_cur, v_idxs, il, sinfos[i_cur]);
|
||||||
}
|
}
|
||||||
|
|
||||||
ggml_tensor * llama_kv_cache_context::cpy_k_idx(ggml_context * ctx, ggml_tensor * k_idx_cur, ggml_tensor * k_idxs, int32_t il) const {
|
|
||||||
return kv->cpy_k_idx(ctx, k_idx_cur, k_idxs, il, sinfos[i_cur]);
|
|
||||||
}
|
|
||||||
|
|
||||||
ggml_tensor * llama_kv_cache_context::build_input_k_idxs(ggml_context * ctx, const llama_ubatch & ubatch) const {
|
ggml_tensor * llama_kv_cache_context::build_input_k_idxs(ggml_context * ctx, const llama_ubatch & ubatch) const {
|
||||||
return kv->build_input_k_idxs(ctx, ubatch);
|
return kv->build_input_k_idxs(ctx, ubatch);
|
||||||
}
|
}
|
||||||
|
|||||||
+2
-10
@@ -164,6 +164,8 @@ public:
|
|||||||
std::vector<uint32_t> get_layer_ids() const;
|
std::vector<uint32_t> get_layer_ids() const;
|
||||||
ggml_tensor * get_k_storage(int32_t il) const;
|
ggml_tensor * get_k_storage(int32_t il) const;
|
||||||
|
|
||||||
|
const llama_kv_cells & get_cells(llama_seq_id seq_id) const;
|
||||||
|
|
||||||
//
|
//
|
||||||
// graph_build API
|
// graph_build API
|
||||||
//
|
//
|
||||||
@@ -173,12 +175,10 @@ public:
|
|||||||
// get views of the current state of the cache
|
// get views of the current state of the cache
|
||||||
ggml_tensor * get_k(ggml_context * ctx, int32_t il, uint32_t n_kv, const slot_info & sinfo) const;
|
ggml_tensor * get_k(ggml_context * ctx, int32_t il, uint32_t n_kv, const slot_info & sinfo) const;
|
||||||
ggml_tensor * get_v(ggml_context * ctx, int32_t il, uint32_t n_kv, const slot_info & sinfo) const;
|
ggml_tensor * get_v(ggml_context * ctx, int32_t il, uint32_t n_kv, const slot_info & sinfo) const;
|
||||||
ggml_tensor * get_k_idx(ggml_context * ctx, int32_t il, uint32_t n_kv, const slot_info & sinfo) const;
|
|
||||||
|
|
||||||
// store k_cur and v_cur in the cache based on the provided head location
|
// store k_cur and v_cur in the cache based on the provided head location
|
||||||
ggml_tensor * cpy_k(ggml_context * ctx, ggml_tensor * k_cur, ggml_tensor * k_idxs, int32_t il, const slot_info & sinfo) const;
|
ggml_tensor * cpy_k(ggml_context * ctx, ggml_tensor * k_cur, ggml_tensor * k_idxs, int32_t il, const slot_info & sinfo) const;
|
||||||
ggml_tensor * cpy_v(ggml_context * ctx, ggml_tensor * v_cur, ggml_tensor * v_idxs, int32_t il, const slot_info & sinfo) const;
|
ggml_tensor * cpy_v(ggml_context * ctx, ggml_tensor * v_cur, ggml_tensor * v_idxs, int32_t il, const slot_info & sinfo) const;
|
||||||
ggml_tensor * cpy_k_idx(ggml_context * ctx, ggml_tensor * k_idx_cur, ggml_tensor * k_idxs, int32_t il, const slot_info & sinfo) const;
|
|
||||||
|
|
||||||
//
|
//
|
||||||
// preparation API
|
// preparation API
|
||||||
@@ -230,11 +230,9 @@ private:
|
|||||||
|
|
||||||
ggml_tensor * k;
|
ggml_tensor * k;
|
||||||
ggml_tensor * v;
|
ggml_tensor * v;
|
||||||
ggml_tensor * k_idx; // MSA single-head indexer keys, F32
|
|
||||||
|
|
||||||
std::vector<ggml_tensor *> k_stream;
|
std::vector<ggml_tensor *> k_stream;
|
||||||
std::vector<ggml_tensor *> v_stream;
|
std::vector<ggml_tensor *> v_stream;
|
||||||
std::vector<ggml_tensor *> k_idx_stream;
|
|
||||||
};
|
};
|
||||||
|
|
||||||
bool v_trans = true; // the value tensor is transposed
|
bool v_trans = true; // the value tensor is transposed
|
||||||
@@ -263,9 +261,6 @@ private:
|
|||||||
// env: LLAMA_KV_CACHE_DEBUG
|
// env: LLAMA_KV_CACHE_DEBUG
|
||||||
int debug = 0;
|
int debug = 0;
|
||||||
|
|
||||||
// set when a k_idx (indexer) cache exists and the stream layout supports MSA (single seq, or one stream per seq)
|
|
||||||
bool msa_strict_slots = false;
|
|
||||||
|
|
||||||
// this is the SWA type of the cache - not to be confused with the model SWA type
|
// this is the SWA type of the cache - not to be confused with the model SWA type
|
||||||
const llama_swa_type swa_type = LLAMA_SWA_TYPE_NONE;
|
const llama_swa_type swa_type = LLAMA_SWA_TYPE_NONE;
|
||||||
|
|
||||||
@@ -298,7 +293,6 @@ private:
|
|||||||
|
|
||||||
size_t size_k_bytes() const;
|
size_t size_k_bytes() const;
|
||||||
size_t size_v_bytes() const;
|
size_t size_v_bytes() const;
|
||||||
size_t size_k_idx_bytes() const;
|
|
||||||
|
|
||||||
ggml_tensor * build_rope_shift(
|
ggml_tensor * build_rope_shift(
|
||||||
const llama_cparams & cparams,
|
const llama_cparams & cparams,
|
||||||
@@ -378,7 +372,6 @@ public:
|
|||||||
// get views of the current state of the cache
|
// get views of the current state of the cache
|
||||||
ggml_tensor * get_k(ggml_context * ctx, int32_t il) const;
|
ggml_tensor * get_k(ggml_context * ctx, int32_t il) const;
|
||||||
ggml_tensor * get_v(ggml_context * ctx, int32_t il) const;
|
ggml_tensor * get_v(ggml_context * ctx, int32_t il) const;
|
||||||
ggml_tensor * get_k_idx(ggml_context * ctx, int32_t il) const;
|
|
||||||
|
|
||||||
// store k_cur and v_cur in the cache based on the provided head location
|
// store k_cur and v_cur in the cache based on the provided head location
|
||||||
// note: the heads in k_cur and v_cur should be laid out contiguously in memory
|
// note: the heads in k_cur and v_cur should be laid out contiguously in memory
|
||||||
@@ -388,7 +381,6 @@ public:
|
|||||||
// - v_idxs [n_tokens] or [n_tokens*n_embd_v_gqa] depending if V cache is transposed
|
// - v_idxs [n_tokens] or [n_tokens*n_embd_v_gqa] depending if V cache is transposed
|
||||||
ggml_tensor * cpy_k(ggml_context * ctx, ggml_tensor * k_cur, ggml_tensor * k_idxs, int32_t il) const;
|
ggml_tensor * cpy_k(ggml_context * ctx, ggml_tensor * k_cur, ggml_tensor * k_idxs, int32_t il) const;
|
||||||
ggml_tensor * cpy_v(ggml_context * ctx, ggml_tensor * v_cur, ggml_tensor * v_idxs, int32_t il) const;
|
ggml_tensor * cpy_v(ggml_context * ctx, ggml_tensor * v_cur, ggml_tensor * v_idxs, int32_t il) const;
|
||||||
ggml_tensor * cpy_k_idx(ggml_context * ctx, ggml_tensor * k_idx_cur, ggml_tensor * k_idxs, int32_t il) const;
|
|
||||||
|
|
||||||
// create destination indices for each head of the current batch for where it would be written in the KV cache
|
// create destination indices for each head of the current batch for where it would be written in the KV cache
|
||||||
// the indices address the global KV cache (not per stream) - this is not relevant for the user of this API, but
|
// the indices address the global KV cache (not per stream) - this is not relevant for the user of this API, but
|
||||||
|
|||||||
@@ -11,6 +11,7 @@
|
|||||||
#include "llama-kv-cache.h"
|
#include "llama-kv-cache.h"
|
||||||
#include "llama-kv-cache-iswa.h"
|
#include "llama-kv-cache-iswa.h"
|
||||||
#include "llama-kv-cache-dsa.h"
|
#include "llama-kv-cache-dsa.h"
|
||||||
|
#include "llama-kv-cache-msa.h"
|
||||||
#include "llama-kv-cache-dsv4.h"
|
#include "llama-kv-cache-dsv4.h"
|
||||||
#include "llama-memory-hybrid.h"
|
#include "llama-memory-hybrid.h"
|
||||||
#include "llama-memory-hybrid-iswa.h"
|
#include "llama-memory-hybrid-iswa.h"
|
||||||
@@ -2071,6 +2072,28 @@ llama_memory_i * llama_model::create_memory(const llama_memory_params & params,
|
|||||||
{
|
{
|
||||||
res = nullptr;
|
res = nullptr;
|
||||||
} break;
|
} break;
|
||||||
|
case LLM_ARCH_MINIMAX_M3:
|
||||||
|
{
|
||||||
|
// sparse (MSA) layers carry an indexer key cache, but leading dense layers do not
|
||||||
|
llama_kv_cache::layer_filter_cb filter_idx =
|
||||||
|
[&](int32_t il) { return (uint32_t) il >= hparams.n_layer_dense_lead; };
|
||||||
|
|
||||||
|
res = new llama_kv_cache_msa(
|
||||||
|
*this,
|
||||||
|
params.type_k,
|
||||||
|
params.type_v,
|
||||||
|
!cparams.flash_attn,
|
||||||
|
cparams.offload_kqv,
|
||||||
|
cparams.kv_unified,
|
||||||
|
cparams.n_ctx_seq,
|
||||||
|
cparams.n_seq_max,
|
||||||
|
1,
|
||||||
|
hparams.n_swa,
|
||||||
|
hparams.swa_type,
|
||||||
|
nullptr,
|
||||||
|
filter_idx,
|
||||||
|
nullptr);
|
||||||
|
} break;
|
||||||
case LLM_ARCH_GLM_DSA:
|
case LLM_ARCH_GLM_DSA:
|
||||||
case LLM_ARCH_DEEPSEEK32:
|
case LLM_ARCH_DEEPSEEK32:
|
||||||
{
|
{
|
||||||
|
|||||||
+152
-75
@@ -1,5 +1,5 @@
|
|||||||
#include "models.h"
|
#include "models.h"
|
||||||
#include "llama-kv-cache.h"
|
#include "llama-kv-cache-msa.h"
|
||||||
#include <cmath>
|
#include <cmath>
|
||||||
#include <vector>
|
#include <vector>
|
||||||
#include <cstdint>
|
#include <cstdint>
|
||||||
@@ -7,7 +7,8 @@
|
|||||||
// MiniMax-M3: MiniMax-M2 style GQA (per-head QK-norm, partial rotary) with
|
// MiniMax-M3: MiniMax-M2 style GQA (per-head QK-norm, partial rotary) with
|
||||||
// DeepSeek-V3 leading-dense + routed/shared experts (sigmoid gating, routed scaling),
|
// DeepSeek-V3 leading-dense + routed/shared experts (sigmoid gating, routed scaling),
|
||||||
// swigluoai activation, and MiniMax Sparse Attention (MSA). MTP is not in released model weights.
|
// swigluoai activation, and MiniMax Sparse Attention (MSA). MTP is not in released model weights.
|
||||||
// Notes: Blocks are anchored to absolute KV cache slots.
|
// MSA blocks are defined over token positions. The graph translates between position space (block
|
||||||
|
// selection) and cell space (K/V/indexer storage) via per-ubatch pos<->cell maps populated from llama_kv_cells
|
||||||
|
|
||||||
void llama_model_minimax_m3::load_arch_hparams(llama_model_loader & ml) {
|
void llama_model_minimax_m3::load_arch_hparams(llama_model_loader & ml) {
|
||||||
ml.get_key(LLM_KV_ATTENTION_LAYERNORM_RMS_EPS, hparams.f_norm_rms_eps);
|
ml.get_key(LLM_KV_ATTENTION_LAYERNORM_RMS_EPS, hparams.f_norm_rms_eps);
|
||||||
@@ -23,7 +24,6 @@ void llama_model_minimax_m3::load_arch_hparams(llama_model_loader & ml) {
|
|||||||
ml.get_key(LLM_KV_ATTENTION_INDEXER_BLOCK_SIZE, hparams.indexer_block_size);
|
ml.get_key(LLM_KV_ATTENTION_INDEXER_BLOCK_SIZE, hparams.indexer_block_size);
|
||||||
ml.get_key(LLM_KV_ATTENTION_INDEXER_LOCAL_BLOCKS, hparams.indexer_local_blocks);
|
ml.get_key(LLM_KV_ATTENTION_INDEXER_LOCAL_BLOCKS, hparams.indexer_local_blocks);
|
||||||
msa_p = { (int) hparams.indexer_block_size, (int) hparams.indexer_top_k, (int) hparams.indexer_local_blocks };
|
msa_p = { (int) hparams.indexer_block_size, (int) hparams.indexer_top_k, (int) hparams.indexer_local_blocks };
|
||||||
hparams.indexer_kv = true;
|
|
||||||
|
|
||||||
switch (hparams.n_layer()) {
|
switch (hparams.n_layer()) {
|
||||||
case 60: type = LLM_TYPE_428B_A23B; break;
|
case 60: type = LLM_TYPE_428B_A23B; break;
|
||||||
@@ -86,43 +86,83 @@ std::unique_ptr<llm_graph_context> llama_model_minimax_m3::build_arch_graph(cons
|
|||||||
return std::make_unique<graph>(*this, params);
|
return std::make_unique<graph>(*this, params);
|
||||||
}
|
}
|
||||||
|
|
||||||
// per-query local-force bias for MSA selection
|
class llm_graph_input_msa : public llm_graph_input_i {
|
||||||
// local window always wins a slot
|
|
||||||
class llm_graph_input_msa_local : public llm_graph_input_i {
|
|
||||||
public:
|
public:
|
||||||
llm_graph_input_msa_local(int blk, int local, int64_t nblk) : blk(blk), local(local), nblk(nblk) {}
|
llm_graph_input_msa(const llama_kv_cache_msa_context * mctx, int blk, int local) :
|
||||||
|
mctx(mctx), blk(blk), local(local) {}
|
||||||
|
|
||||||
void set_input(const llama_ubatch * ubatch) override {
|
void set_input(const llama_ubatch * ubatch) override {
|
||||||
if (!bias || !ubatch->pos) {
|
if (pos_slot_i) { mctx->set_input_pos_slot(pos_slot_i, ubatch); }
|
||||||
return;
|
if (pos_slot_f) { mctx->set_input_pos_slot(pos_slot_f, ubatch); }
|
||||||
}
|
if (cell_blk) { mctx->set_input_cell_pos(cell_blk, ubatch, blk); }
|
||||||
const int64_t n_tokens = ubatch->n_tokens;
|
if (pos_mask) { mctx->set_input_pos_mask(pos_mask, ubatch); }
|
||||||
std::vector<float> data((size_t) nblk * n_tokens, 0.0f);
|
|
||||||
for (int64_t i = 0; i < n_tokens; ++i) {
|
// local-force bias over position blocks
|
||||||
const int64_t L = ubatch->pos[i] / blk;
|
if (bias && ubatch->pos) {
|
||||||
for (int l = 0; l < local && L - l >= 0; ++l) {
|
const int64_t n_tokens = ubatch->n_tokens;
|
||||||
if (L - l < nblk) {
|
const int64_t nblk = bias->ne[0];
|
||||||
data[(size_t) i * nblk + (L - l)] = 1e30f;
|
std::vector<float> data((size_t) nblk * n_tokens, 0.0f);
|
||||||
|
for (int64_t i = 0; i < n_tokens; ++i) {
|
||||||
|
const int64_t L = ubatch->pos[i] / blk;
|
||||||
|
for (int l = 0; l < local && L - l >= 0; ++l) {
|
||||||
|
if (L - l < nblk) {
|
||||||
|
data[(size_t) i * nblk + (L - l)] = 1e30f;
|
||||||
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
ggml_backend_tensor_set(bias, data.data(), 0, data.size() * sizeof(float));
|
||||||
}
|
}
|
||||||
ggml_backend_tensor_set(bias, data.data(), 0, data.size() * sizeof(float));
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// valid as long as the bias tensor dims still match the new ubatch/cache window
|
// valid as long as the tensor dims still match the new ubatch/cache window and the
|
||||||
|
// ubatch is in the same regime (decode graphs have pos_slot_f, batch graphs cell_blk)
|
||||||
bool can_reuse(const llm_graph_params & params) override {
|
bool can_reuse(const llm_graph_params & params) override {
|
||||||
const auto * mctx = static_cast<const llama_kv_cache_context *>(params.mctx);
|
const auto * mctx_new = static_cast<const llama_kv_cache_msa_context *>(params.mctx);
|
||||||
|
|
||||||
|
this->mctx = mctx_new;
|
||||||
|
|
||||||
|
const int64_t n_ps = GGML_PAD((int64_t) mctx_new->get_n_pos(), blk);
|
||||||
|
const int64_t ns = params.cparams.kv_unified ? 1 : params.ubatch.n_seqs_unq;
|
||||||
|
|
||||||
|
const bool decode = params.ubatch.n_tokens == ns; // one token per stream
|
||||||
|
|
||||||
bool res = true;
|
bool res = true;
|
||||||
res &= bias->ne[1] == params.ubatch.n_tokens;
|
|
||||||
res &= bias->ne[0] * blk == (int64_t) mctx->get_n_kv();
|
res &= bias->ne[0] * blk == n_ps;
|
||||||
|
res &= bias->ne[1] == params.ubatch.n_tokens;
|
||||||
|
|
||||||
|
res &= pos_mask->ne[0] == n_ps;
|
||||||
|
res &= pos_mask->ne[1] == params.ubatch.n_tokens;
|
||||||
|
|
||||||
|
res &= pos_slot_i->ne[0] == n_ps;
|
||||||
|
res &= pos_slot_i->ne[1] == ns;
|
||||||
|
|
||||||
|
res &= decode == (pos_slot_f != nullptr);
|
||||||
|
res &= decode == (cell_blk == nullptr);
|
||||||
|
|
||||||
|
if (pos_slot_f) {
|
||||||
|
res &= pos_slot_f->ne[0] == n_ps;
|
||||||
|
res &= pos_slot_f->ne[1] == ns;
|
||||||
|
}
|
||||||
|
|
||||||
|
if (cell_blk) {
|
||||||
|
res &= cell_blk->ne[0] == (int64_t) mctx_new->get_base()->get_n_kv();
|
||||||
|
res &= cell_blk->ne[1] == ns;
|
||||||
|
}
|
||||||
|
|
||||||
return res;
|
return res;
|
||||||
}
|
}
|
||||||
|
|
||||||
ggml_tensor * bias = nullptr;
|
ggml_tensor * bias = nullptr; // F32 [nblk, n_tokens] local-force bias (position blocks)
|
||||||
int blk;
|
ggml_tensor * pos_mask = nullptr; // F32 [n_ps, n_tokens] 0/-inf visibility, by position
|
||||||
int local;
|
ggml_tensor * pos_slot_i = nullptr; // I32 [n_ps, ns] pos -> cell (get_rows index)
|
||||||
int64_t nblk;
|
ggml_tensor * pos_slot_f = nullptr; // F32 [n_ps, ns] pos -> cell (gatherable values, decode)
|
||||||
|
ggml_tensor * cell_blk = nullptr; // I32 [n_kv, ns] cell -> position block (batch)
|
||||||
|
|
||||||
|
const llama_kv_cache_msa_context * mctx;
|
||||||
|
|
||||||
|
int blk;
|
||||||
|
int local;
|
||||||
};
|
};
|
||||||
|
|
||||||
// One FA call for all GQA groups (and at multi-stream decode, all streams) by mapping them onto the FA sequence dim (ne[3])
|
// One FA call for all GQA groups (and at multi-stream decode, all streams) by mapping them onto the FA sequence dim (ne[3])
|
||||||
@@ -173,7 +213,7 @@ llama_model_minimax_m3::graph::graph(const llama_model & model, const llm_graph_
|
|||||||
inpL = build_inp_embd(model.tok_embd);
|
inpL = build_inp_embd(model.tok_embd);
|
||||||
|
|
||||||
ggml_tensor * inp_pos = build_inp_pos();
|
ggml_tensor * inp_pos = build_inp_pos();
|
||||||
auto inp_attn = build_attn_inp_kv();
|
auto inp_attn = build_attn_inp_kv_msa();
|
||||||
|
|
||||||
// MSA calls ggml_flash_attn_ext directly and assumes the non-transposed V layout that
|
// MSA calls ggml_flash_attn_ext directly and assumes the non-transposed V layout that
|
||||||
// llama.cpp only provides when flash attention is enabled. Block selection is anchored
|
// llama.cpp only provides when flash attention is enabled. Block selection is anchored
|
||||||
@@ -199,34 +239,51 @@ llama_model_minimax_m3::graph::graph(const llama_model & model, const llm_graph_
|
|||||||
}
|
}
|
||||||
|
|
||||||
// hoisted per-graph MSA state (shared by every sparse layer)
|
// hoisted per-graph MSA state (shared by every sparse layer)
|
||||||
llm_graph_input_msa_local * msa_loc = nullptr;
|
llm_graph_input_msa * msa = nullptr;
|
||||||
ggml_tensor * msa_kqm = nullptr;
|
ggml_tensor * msa_kqm = nullptr;
|
||||||
ggml_tensor * msa_mf = nullptr;
|
ggml_tensor * msa_mf = nullptr; // F32 copy of the FA mask for the final mask add
|
||||||
int64_t n_kv = 0, nblk = 0, ns = 1, n_tps = 0;
|
int64_t n_kv = 0, n_ps = 0, nblk = 0, ns = 1, n_tps = 0;
|
||||||
bool msa_decode = false; // gather (1 token per stream) vs mask
|
bool msa_decode = false; // gather (1 token per stream) vs mask
|
||||||
const int blk = mm.msa_p.blk;
|
const int blk = mm.msa_p.blk;
|
||||||
const int64_t Hd = hparams.indexer_n_head; // one indexer head per GQA group
|
const int64_t Hd = hparams.indexer_n_head; // one indexer head per GQA group
|
||||||
|
|
||||||
if (msa_enabled) {
|
if (msa_enabled) {
|
||||||
|
const auto * mctx_msa = static_cast<const llama_kv_cache_msa_context *>(mctx);
|
||||||
|
|
||||||
msa_kqm = inp_attn->get_kq_mask();
|
msa_kqm = inp_attn->get_kq_mask();
|
||||||
n_kv = msa_kqm->ne[0];
|
n_kv = msa_kqm->ne[0];
|
||||||
n_tps = msa_kqm->ne[1]; // tokens per stream
|
n_tps = msa_kqm->ne[1]; // tokens per stream
|
||||||
ns = msa_kqm->ne[3]; // streams in this ubatch
|
ns = msa_kqm->ne[3]; // streams in this ubatch
|
||||||
GGML_ASSERT(msa_kqm->type == GGML_TYPE_F16 && "MSA requires the FA (f16) mask");
|
GGML_ASSERT(msa_kqm->type == GGML_TYPE_F16 && "MSA requires the FA (f16) mask");
|
||||||
GGML_ASSERT(n_tps*ns == n_tokens);
|
GGML_ASSERT(n_tps*ns == n_tokens);
|
||||||
GGML_ASSERT(n_kv % blk == 0 &&
|
|
||||||
"MSA: KV/mask n_kv must be a multiple of indexer.block_size (128); "
|
// the position axis covers every position currently in the cache and is padded to whole blocks
|
||||||
"the flash-attention KV padding must be a multiple of the block size. "
|
n_ps = GGML_PAD((int64_t) mctx_msa->get_n_pos(), blk);
|
||||||
"A non-multiple would silently drop the partial tail block.");
|
nblk = n_ps / blk;
|
||||||
nblk = n_kv / blk;
|
|
||||||
msa_decode = n_tps == 1;
|
msa_decode = n_tps == 1;
|
||||||
|
|
||||||
msa_mf = ggml_cast(ctx0, msa_kqm, GGML_TYPE_F32);
|
auto inp = std::make_unique<llm_graph_input_msa>(mctx_msa, blk, mm.msa_p.local);
|
||||||
|
|
||||||
auto loc = std::make_unique<llm_graph_input_msa_local>(blk, mm.msa_p.local, nblk);
|
inp->bias = ggml_new_tensor_2d(ctx0, GGML_TYPE_F32, nblk, n_tokens); // stream-grouped tokens
|
||||||
loc->bias = ggml_new_tensor_2d(ctx0, GGML_TYPE_F32, nblk, n_tokens); // stream-grouped tokens
|
ggml_set_input(inp->bias);
|
||||||
ggml_set_input(loc->bias);
|
|
||||||
msa_loc = (llm_graph_input_msa_local *) res->add_input(std::move(loc));
|
inp->pos_mask = ggml_new_tensor_2d(ctx0, GGML_TYPE_F32, n_ps, n_tokens);
|
||||||
|
ggml_set_input(inp->pos_mask);
|
||||||
|
|
||||||
|
inp->pos_slot_i = ggml_new_tensor_2d(ctx0, GGML_TYPE_I32, n_ps, ns);
|
||||||
|
ggml_set_input(inp->pos_slot_i);
|
||||||
|
|
||||||
|
if (msa_decode) {
|
||||||
|
inp->pos_slot_f = ggml_new_tensor_2d(ctx0, GGML_TYPE_F32, n_ps, ns);
|
||||||
|
ggml_set_input(inp->pos_slot_f);
|
||||||
|
} else {
|
||||||
|
inp->cell_blk = ggml_new_tensor_2d(ctx0, GGML_TYPE_I32, n_kv, ns);
|
||||||
|
ggml_set_input(inp->cell_blk);
|
||||||
|
|
||||||
|
msa_mf = ggml_cast(ctx0, msa_kqm, GGML_TYPE_F32);
|
||||||
|
}
|
||||||
|
|
||||||
|
msa = (llm_graph_input_msa *) res->add_input(std::move(inp));
|
||||||
}
|
}
|
||||||
|
|
||||||
ggml_tensor * inp_out_ids = build_inp_out_ids();
|
ggml_tensor * inp_out_ids = build_inp_out_ids();
|
||||||
@@ -283,9 +340,11 @@ llama_model_minimax_m3::graph::graph(const llama_model & model, const llm_graph_
|
|||||||
ik = ggml_rope_ext(ctx0, ik, inp_pos, nullptr, n_rot, rope_type, n_ctx_orig,
|
ik = ggml_rope_ext(ctx0, ik, inp_pos, nullptr, n_rot, rope_type, n_ctx_orig,
|
||||||
freq_base, freq_scale, ext_factor, attn_factor, beta_fast, beta_slow);
|
freq_base, freq_scale, ext_factor, attn_factor, beta_fast, beta_slow);
|
||||||
|
|
||||||
const auto * mctx_cur = inp_attn->mctx;
|
const auto * mctx_msa_l = static_cast<const llama_kv_cache_msa_context *>(mctx);
|
||||||
ggml_build_forward_expand(gf, mctx_cur->cpy_k_idx(ctx0, ik, inp_attn->get_k_idxs(), il));
|
const auto * mctx_cur = mctx_msa_l->get_base();
|
||||||
ggml_tensor * ik_kv = mctx_cur->get_k_idx(ctx0, il);
|
const auto * mctx_idx = mctx_msa_l->get_idx();
|
||||||
|
ggml_build_forward_expand(gf, mctx_idx->cpy_k(ctx0, ik, inp_attn->get_k_idxs_idx(), il));
|
||||||
|
ggml_tensor * ik_kv = mctx_idx->get_k(ctx0, il);
|
||||||
|
|
||||||
if (inp_attn->self_k_rot) {
|
if (inp_attn->self_k_rot) {
|
||||||
Qcur = llama_mul_mat_hadamard(ctx0, Qcur, inp_attn->self_k_rot);
|
Qcur = llama_mul_mat_hadamard(ctx0, Qcur, inp_attn->self_k_rot);
|
||||||
@@ -316,42 +375,52 @@ llama_model_minimax_m3::graph::graph(const llama_model & model, const llm_graph_
|
|||||||
|
|
||||||
if (msa_decode) {
|
if (msa_decode) {
|
||||||
// decode: batched over streams top-k + gather, one grouped FA
|
// decode: batched over streams top-k + gather, one grouped FA
|
||||||
// scores: per-stream batched matmul over the stream dim (ne[3]).
|
// gather the indexer keys through the pos -> cell map
|
||||||
// the cache views are not contiguous across streams (stride = kv_size, not n_kv)
|
ggml_tensor * ik3 = ggml_view_3d(ctx0, ik_kv, n_idx_dim, n_kv, ns,
|
||||||
ggml_tensor * ikv4 = ggml_view_4d(ctx0, ik_kv, n_idx_dim, n_kv, 1, ns,
|
ik_kv->nb[2], ik_kv->nb[3], 0);
|
||||||
ik_kv->nb[2], ik_kv->nb[3], ik_kv->nb[3], 0);
|
ggml_tensor * ikp = ggml_get_rows(ctx0, ik3, msa->pos_slot_i); // [n_idx_dim, n_ps, ns]
|
||||||
ggml_tensor * iq4 = ggml_reshape_4d(ctx0, iq, n_idx_dim, Hd, 1, ns);
|
ggml_tensor * iq4 = ggml_reshape_4d(ctx0, iq, n_idx_dim, Hd, 1, ns);
|
||||||
ggml_tensor * sc = ggml_mul_mat(ctx0, ikv4, iq4);
|
ggml_tensor * sc = ggml_mul_mat(ctx0,
|
||||||
|
ggml_reshape_4d(ctx0, ikp, n_idx_dim, n_ps, 1, ns), iq4);
|
||||||
ggml_mul_mat_set_prec(sc, GGML_PREC_F32);
|
ggml_mul_mat_set_prec(sc, GGML_PREC_F32);
|
||||||
sc = ggml_add_inplace(ctx0, sc, msa_mf);
|
// unmapped positions come out -inf, so they can never rank into the top-k
|
||||||
|
sc = ggml_add_inplace(ctx0, sc,
|
||||||
|
ggml_reshape_4d(ctx0, msa->pos_mask, n_ps, 1, 1, ns));
|
||||||
ggml_tensor * bs = ggml_pool_2d(ctx0, sc, GGML_OP_POOL_MAX, blk, 1, blk, 1, 0, 0);
|
ggml_tensor * bs = ggml_pool_2d(ctx0, sc, GGML_OP_POOL_MAX, blk, 1, blk, 1, 0, 0);
|
||||||
cb(bs, "msa_bs", il);
|
cb(bs, "msa_bs", il);
|
||||||
|
|
||||||
ggml_tensor * bsf = ggml_add(ctx0, bs,
|
ggml_tensor * bsf = ggml_add(ctx0, bs,
|
||||||
ggml_reshape_4d(ctx0, msa_loc->bias, nblk, 1, 1, ns));
|
ggml_reshape_4d(ctx0, msa->bias, nblk, 1, 1, ns));
|
||||||
ggml_tensor * idx = ggml_top_k(ctx0, bsf, K);
|
ggml_tensor * idx = ggml_top_k(ctx0, bsf, K); // position blocks
|
||||||
|
|
||||||
// token idx: tj[t,k,h,s] = blk*idx[k,h,s] + t (for the mask gather)
|
// pos idx: tj[t,k,h,s] = blk*idx[k,h,s] + t (positions - mask gather)
|
||||||
// row idx: tr[t,k,h,s] = tj*HKV + h (for the per-stream K/V gather)
|
// cell idx: cs[t,k,h,s] = pos_slot[tj] (pos -> cell translation)
|
||||||
|
// row idx: tr[t,k,h,s] = cs*HKV + h (per-stream K/V gather)
|
||||||
ggml_tensor * a = ggml_scale(ctx0, ggml_cast(ctx0, idx, GGML_TYPE_F32), (float) blk);
|
ggml_tensor * a = ggml_scale(ctx0, ggml_cast(ctx0, idx, GGML_TYPE_F32), (float) blk);
|
||||||
a = ggml_reshape_4d(ctx0, a, 1, K, Hd, ns);
|
a = ggml_reshape_4d(ctx0, a, 1, K, Hd, ns);
|
||||||
ggml_tensor * tj = ggml_add(ctx0,
|
ggml_tensor * tj = ggml_add(ctx0,
|
||||||
ggml_repeat_4d(ctx0, a, blk, K, Hd, ns),
|
ggml_repeat_4d(ctx0, a, blk, K, Hd, ns),
|
||||||
ggml_reshape_3d(ctx0, ggml_arange(ctx0, 0.0f, (float) blk, 1.0f), blk, 1, 1));
|
ggml_reshape_3d(ctx0, ggml_arange(ctx0, 0.0f, (float) blk, 1.0f), blk, 1, 1));
|
||||||
ggml_tensor * tr = ggml_add(ctx0,
|
|
||||||
ggml_scale(ctx0, tj, (float) HKV),
|
|
||||||
ggml_reshape_3d(ctx0, ggml_arange(ctx0, 0.0f, (float) HKV, 1.0f), 1, 1, Hd));
|
|
||||||
|
|
||||||
ggml_tensor * tokj = ggml_cast(ctx0, ggml_reshape_2d(ctx0, tj, (int64_t) blk*K*Hd, ns), GGML_TYPE_I32);
|
ggml_tensor * tokj = ggml_cast(ctx0, ggml_reshape_2d(ctx0, tj, (int64_t) blk*K*Hd, ns), GGML_TYPE_I32);
|
||||||
|
|
||||||
|
ggml_tensor * cs = ggml_get_rows(ctx0,
|
||||||
|
ggml_reshape_3d(ctx0, msa->pos_slot_f, 1, n_ps, ns), tokj); // [1, blk*K*Hd, ns]
|
||||||
|
cs = ggml_reshape_4d(ctx0, cs, blk, K, Hd, ns);
|
||||||
|
|
||||||
|
ggml_tensor * tr = ggml_add(ctx0,
|
||||||
|
ggml_scale(ctx0, cs, (float) HKV),
|
||||||
|
ggml_reshape_3d(ctx0, ggml_arange(ctx0, 0.0f, (float) HKV, 1.0f), 1, 1, Hd));
|
||||||
|
|
||||||
ggml_tensor * tokr = ggml_cast(ctx0, ggml_reshape_2d(ctx0, tr, (int64_t) blk*K*Hd, ns), GGML_TYPE_I32);
|
ggml_tensor * tokr = ggml_cast(ctx0, ggml_reshape_2d(ctx0, tr, (int64_t) blk*K*Hd, ns), GGML_TYPE_I32);
|
||||||
|
|
||||||
ggml_tensor * k3 = ggml_view_3d(ctx0, k, D, HKV*n_kv, ns, k->nb[1], k->nb[3], 0);
|
ggml_tensor * k3 = ggml_view_3d(ctx0, k, D, HKV*n_kv, ns, k->nb[1], k->nb[3], 0);
|
||||||
ggml_tensor * v3 = ggml_view_3d(ctx0, v, D, HKV*n_kv, ns, v->nb[1], v->nb[3], 0);
|
ggml_tensor * v3 = ggml_view_3d(ctx0, v, D, HKV*n_kv, ns, v->nb[1], v->nb[3], 0);
|
||||||
ggml_tensor * m3 = ggml_reshape_3d(ctx0, msa_kqm, 1, n_kv, ns);
|
ggml_tensor * mp = ggml_reshape_3d(ctx0, msa->pos_mask, 1, n_ps, ns);
|
||||||
|
|
||||||
ggml_tensor * kg = ggml_get_rows(ctx0, k3, tokr);
|
ggml_tensor * kg = ggml_get_rows(ctx0, k3, tokr);
|
||||||
ggml_tensor * vg = ggml_get_rows(ctx0, v3, tokr);
|
ggml_tensor * vg = ggml_get_rows(ctx0, v3, tokr);
|
||||||
ggml_tensor * mg = ggml_get_rows(ctx0, m3, tokj);
|
ggml_tensor * mg = ggml_get_rows(ctx0, mp, tokj);
|
||||||
|
|
||||||
// fold (group, stream) onto the FA channel dim
|
// fold (group, stream) onto the FA channel dim
|
||||||
const ggml_type kt = ggml_is_quantized(k->type) ? GGML_TYPE_F16 : k->type;
|
const ggml_type kt = ggml_is_quantized(k->type) ? GGML_TYPE_F16 : k->type;
|
||||||
@@ -372,12 +441,16 @@ llama_model_minimax_m3::graph::graph(const llama_model & model, const llm_graph_
|
|||||||
iq->nb[1], iq->nb[2], st*n_tps*iq->nb[2]);
|
iq->nb[1], iq->nb[2], st*n_tps*iq->nb[2]);
|
||||||
ggml_tensor * ik_s = ggml_view_2d(ctx0, ik_kv, n_idx_dim, n_kv,
|
ggml_tensor * ik_s = ggml_view_2d(ctx0, ik_kv, n_idx_dim, n_kv,
|
||||||
ik_kv->nb[2], st*ik_kv->nb[3]);
|
ik_kv->nb[2], st*ik_kv->nb[3]);
|
||||||
ggml_tensor * mf_s = ggml_view_3d(ctx0, msa_mf, n_kv, 1, n_tps,
|
ggml_tensor * psl_s = ggml_view_1d(ctx0, msa->pos_slot_i, n_ps,
|
||||||
msa_mf->nb[1], msa_mf->nb[1], st*msa_mf->nb[3]);
|
st*msa->pos_slot_i->nb[1]);
|
||||||
ggml_tensor * km_s = ggml_view_3d(ctx0, msa_kqm, n_kv, n_tps, 1,
|
ggml_tensor * pm_s = ggml_view_3d(ctx0, msa->pos_mask, n_ps, 1, n_tps,
|
||||||
msa_kqm->nb[1], msa_kqm->nb[3], st*msa_kqm->nb[3]);
|
msa->pos_mask->nb[1], msa->pos_mask->nb[1], st*n_tps*msa->pos_mask->nb[1]);
|
||||||
ggml_tensor * bias_s = ggml_view_3d(ctx0, msa_loc->bias, nblk, 1, n_tps,
|
ggml_tensor * cb_s = ggml_view_1d(ctx0, msa->cell_blk, n_kv,
|
||||||
msa_loc->bias->nb[1], msa_loc->bias->nb[1], st*n_tps*msa_loc->bias->nb[1]);
|
st*msa->cell_blk->nb[1]);
|
||||||
|
ggml_tensor * mf_s = ggml_view_3d(ctx0, msa_mf, n_kv, n_tps, 1,
|
||||||
|
msa_mf->nb[1], msa_mf->nb[3], st*msa_mf->nb[3]);
|
||||||
|
ggml_tensor * bias_s = ggml_view_3d(ctx0, msa->bias, nblk, 1, n_tps,
|
||||||
|
msa->bias->nb[1], msa->bias->nb[1], st*n_tps*msa->bias->nb[1]);
|
||||||
ggml_tensor * q_s = ggml_view_3d(ctx0, Qcur, D, n_head, n_tps,
|
ggml_tensor * q_s = ggml_view_3d(ctx0, Qcur, D, n_head, n_tps,
|
||||||
Qcur->nb[1], Qcur->nb[2], st*n_tps*Qcur->nb[2]);
|
Qcur->nb[1], Qcur->nb[2], st*n_tps*Qcur->nb[2]);
|
||||||
ggml_tensor * k_s = ggml_view_4d(ctx0, k, D, HKV, n_kv, 1,
|
ggml_tensor * k_s = ggml_view_4d(ctx0, k, D, HKV, n_kv, 1,
|
||||||
@@ -385,14 +458,16 @@ llama_model_minimax_m3::graph::graph(const llama_model & model, const llm_graph_
|
|||||||
ggml_tensor * v_s = ggml_view_4d(ctx0, v, D, HKV, n_kv, 1,
|
ggml_tensor * v_s = ggml_view_4d(ctx0, v, D, HKV, n_kv, 1,
|
||||||
v->nb[1], v->nb[2], v->nb[3], st*v->nb[3]);
|
v->nb[1], v->nb[2], v->nb[3], st*v->nb[3]);
|
||||||
|
|
||||||
// block scores: bs = maxpool_blk(idx_q * idx_k^T + causal mask)
|
// block scores: the indexer keys are gathered through the pos -> cell map first
|
||||||
// scores are unscaled, only the top-k ordering matters
|
// scores are unscaled, only the top-k ordering matters
|
||||||
ggml_tensor * sc = ggml_mul_mat(ctx0, ik_s,
|
ggml_tensor * ikp = ggml_get_rows(ctx0, ik_s, psl_s); // [n_idx_dim, n_ps]
|
||||||
|
ggml_tensor * sc = ggml_mul_mat(ctx0, ikp,
|
||||||
ggml_reshape_2d(ctx0, iq_s, n_idx_dim, Hd*n_tps));
|
ggml_reshape_2d(ctx0, iq_s, n_idx_dim, Hd*n_tps));
|
||||||
// indexer scores run in F32
|
// indexer scores run in F32
|
||||||
ggml_mul_mat_set_prec(sc, GGML_PREC_F32);
|
ggml_mul_mat_set_prec(sc, GGML_PREC_F32);
|
||||||
sc = ggml_reshape_3d(ctx0, sc, n_kv, Hd, n_tps);
|
sc = ggml_reshape_3d(ctx0, sc, n_ps, Hd, n_tps);
|
||||||
sc = ggml_add_inplace(ctx0, sc, mf_s);
|
// unmapped positions (holes, padding, empty cells) come out -inf
|
||||||
|
sc = ggml_add_inplace(ctx0, sc, pm_s);
|
||||||
ggml_tensor * bs = ggml_pool_2d(ctx0, sc, GGML_OP_POOL_MAX, blk, 1, blk, 1, 0, 0);
|
ggml_tensor * bs = ggml_pool_2d(ctx0, sc, GGML_OP_POOL_MAX, blk, 1, blk, 1, 0, 0);
|
||||||
cb(bs, "msa_bs", il);
|
cb(bs, "msa_bs", il);
|
||||||
|
|
||||||
@@ -416,14 +491,16 @@ llama_model_minimax_m3::graph::graph(const llama_model & model, const llm_graph_
|
|||||||
bm = ggml_cont(ctx0, ggml_permute(ctx0, bm, 0, 2, 1, 3)); // [nblk, n_tps, Hd]
|
bm = ggml_cont(ctx0, ggml_permute(ctx0, bm, 0, 2, 1, 3)); // [nblk, n_tps, Hd]
|
||||||
cb(bm, "msa_block_mask", il);
|
cb(bm, "msa_block_mask", il);
|
||||||
|
|
||||||
// expand block -> token granularity (j = bk*blk + t),
|
// expand block -> cell granularity through the cell -> position block
|
||||||
// then combine with the causal mask in place
|
// map, then combine with the causal mask. empty cells are masked by the causal mask.
|
||||||
ggml_tensor * bmx = ggml_repeat_4d(ctx0,
|
ggml_tensor * bm2 = ggml_cont(ctx0, ggml_transpose(ctx0,
|
||||||
ggml_reshape_3d(ctx0, bm, 1, nblk, n_tps*Hd),
|
ggml_reshape_2d(ctx0, bm, nblk, n_tps*Hd))); // [n_tps*Hd, nblk]
|
||||||
blk, nblk, n_tps*Hd, 1);
|
ggml_tensor * bmc = ggml_get_rows(ctx0, bm2, cb_s); // [n_tps*Hd, n_kv] F32
|
||||||
|
ggml_tensor * bmx = ggml_cont(ctx0, ggml_transpose(ctx0, bmc));
|
||||||
bmx = ggml_reshape_3d(ctx0, bmx, n_kv, n_tps, Hd);
|
bmx = ggml_reshape_3d(ctx0, bmx, n_kv, n_tps, Hd);
|
||||||
ggml_tensor * mask4 = ggml_add_inplace(ctx0, bmx, km_s);
|
ggml_tensor * mask4 = ggml_add_inplace(ctx0, bmx, mf_s);
|
||||||
mask4 = ggml_reshape_4d(ctx0, mask4, n_kv, n_tps, 1, Hd);
|
mask4 = ggml_cast(ctx0,
|
||||||
|
ggml_reshape_4d(ctx0, mask4, n_kv, n_tps, 1, Hd), GGML_TYPE_F16);
|
||||||
cb(mask4, "msa_mask4", il);
|
cb(mask4, "msa_mask4", il);
|
||||||
|
|
||||||
// cache views with groups on ne[3];
|
// cache views with groups on ne[3];
|
||||||
|
|||||||
Reference in New Issue
Block a user