mirror of
https://github.com/ggml-org/llama.cpp.git
synced 2026-10-01 18:37:28 -05:00
llama: fix qwen4exp (#29751)
* llama: fix qwen4exp * qwen4exp: keep kq_mask input the same shape
This commit is contained in:
@@ -284,6 +284,10 @@ struct llama_hparams {
|
||||
uint32_t indexer_top_k = 0;
|
||||
uint32_t indexer_kpool = 0; // k-pool size
|
||||
bool indexer_kpool_select_tail = true;
|
||||
// head-size slots per cached indexer row, the last one holds the pooled key
|
||||
uint32_t indexer_kpool_row = 3;
|
||||
// pools are consecutive cells in sequence order, not runs of consecutive positions
|
||||
bool indexer_kpool_by_order = false;
|
||||
// MSA
|
||||
uint32_t indexer_block_size = 0;
|
||||
uint32_t indexer_local_blocks = 0;
|
||||
|
||||
+129
-375
@@ -53,8 +53,9 @@ llama_memory_hybrid_idx::llama_memory_hybrid_idx(
|
||||
mem_idx(filter_idx == nullptr ? nullptr : [&] {
|
||||
// MQA with a single key head of indexer_head_size, as llama_kv_cache_dsa shapes its own
|
||||
std::fill(hparams_idx.n_head_kv_arr.begin(), hparams_idx.n_head_kv_arr.end(), 1);
|
||||
// The glm5 next indexer caches key, gate and pooled values per token
|
||||
hparams_idx.n_embd_head_k_full = model.hparams.indexer_head_size * (model.hparams.indexer_kpool > 0 ? 3 : 1);
|
||||
// a k-pool indexer caches its per-token rows and the pooled key side by side
|
||||
// (glm5-next: key | gate | pooled, qwen4exp: key | pooled)
|
||||
hparams_idx.n_embd_head_k_full = model.hparams.indexer_head_size * (model.hparams.indexer_kpool > 0 ? model.hparams.indexer_kpool_row : 1);
|
||||
|
||||
// the cached indexer keys are raw, rotation happens after pooling at read time, so a
|
||||
// K-shift must not rotate them while the stream copies in the same update still apply
|
||||
@@ -331,328 +332,6 @@ llama_kv_cache * llama_memory_hybrid_idx::get_mem_idx() const {
|
||||
return mem_idx.get();
|
||||
}
|
||||
|
||||
void llama_memory_hybrid_idx::set_input_qsa(
|
||||
ggml_tensor * cell_blk,
|
||||
ggml_tensor * blk_cells,
|
||||
ggml_tensor * blk_pos,
|
||||
ggml_tensor * bias,
|
||||
const llama_ubatch * ubatch,
|
||||
uint32_t ratio,
|
||||
bool blk_bias,
|
||||
bool causal_attn) const {
|
||||
GGML_ASSERT(ratio > 0);
|
||||
GGML_ASSERT(get_mem_idx() != nullptr);
|
||||
|
||||
GGML_ASSERT(ggml_backend_buffer_is_host(cell_blk->buffer));
|
||||
|
||||
const int64_t n_kv = cell_blk->ne[0];
|
||||
const int64_t n_ns = cell_blk->ne[1]; // streams in this ubatch
|
||||
const int64_t n_blocks = blk_pos->ne[0]/(4*n_ns);
|
||||
const int64_t n_tokens = ubatch->n_tokens;
|
||||
const int64_t r = ratio;
|
||||
|
||||
GGML_ASSERT(n_tokens % n_ns == 0);
|
||||
const int64_t n_tps = n_tokens/n_ns; // tokens per stream
|
||||
|
||||
int32_t * dst_cell_blk = (int32_t *) cell_blk->data;
|
||||
int32_t * dst_blk_cells = (int32_t *) blk_cells->data;
|
||||
int32_t * dst_blk_pos = (int32_t *) blk_pos->data;
|
||||
float * dst_bias = (float *) bias->data;
|
||||
|
||||
// a block is keyed on (sequence set, index bucket): a unified cache counts every sequence
|
||||
// from zero, so the bucket alone would pool two sequences into one block
|
||||
GGML_ASSERT(r <= 64);
|
||||
const uint64_t slots_full = r == 64 ? ~uint64_t(0) : ((uint64_t(1) << r) - 1);
|
||||
|
||||
// TODO: this runs per ubatch and is O(n_kv) per stream, about 865 us at 33k context. the cost
|
||||
// is the per-cell scan rather than these allocations, so hoisting them buys nothing
|
||||
std::vector<int32_t> blk_of(n_kv);
|
||||
std::vector<int32_t> cell_grp(n_kv);
|
||||
std::vector<int32_t> grp_head(n_blocks);
|
||||
std::vector<int32_t> grp_next;
|
||||
std::vector<int32_t> grp_first;
|
||||
std::vector<int32_t> grp_slot0;
|
||||
std::vector<uint64_t> grp_slots;
|
||||
std::vector<int32_t> grp_bid;
|
||||
std::vector<int32_t> bid_idx;
|
||||
std::vector<int32_t> bid_cell;
|
||||
std::vector<int32_t> bid_slot0;
|
||||
|
||||
std::vector<int32_t> order;
|
||||
std::vector<int32_t> rank;
|
||||
|
||||
std::fill(dst_blk_pos, dst_blk_pos + 4*n_blocks*n_ns, 0);
|
||||
|
||||
for (int64_t s = 0; s < n_ns; ++s) {
|
||||
// ubatch index s*n_tps belongs to this stream; ask which cells array it uses
|
||||
const llama_seq_id seq_of_stream = ubatch->seq_id[s*n_tps][0];
|
||||
const auto & cells = get_mem_idx()->get_cells(seq_of_stream);
|
||||
|
||||
int32_t * cur_cell_blk = dst_cell_blk + s*n_kv;
|
||||
int32_t * cur_blk_cells = dst_blk_cells + s*(r*n_blocks);
|
||||
|
||||
std::fill(cur_blk_cells, cur_blk_cells + r*n_blocks, 0);
|
||||
|
||||
bid_idx .clear();
|
||||
bid_cell .clear();
|
||||
bid_slot0.clear();
|
||||
|
||||
int n_seq_present = 0;
|
||||
|
||||
for (int sq = 0; sq < LLAMA_MAX_SEQ && n_seq_present < 2; ++sq) {
|
||||
if (cells.seq_pos_min(sq) >= 0) {
|
||||
n_seq_present++;
|
||||
}
|
||||
}
|
||||
|
||||
const bool one_seq = n_seq_present <= 1;
|
||||
|
||||
// a cell no block covers needs its own -inf, which a per-block bias cannot carry
|
||||
// every cache path keeps the position below the cell window, so this stays false
|
||||
bool oor = false;
|
||||
|
||||
bool dup = false;
|
||||
|
||||
bool ranked = false;
|
||||
|
||||
auto group_cells = [&]() {
|
||||
// -1 means no usable block: an incomplete or short group cannot be pooled
|
||||
std::fill(blk_of.begin(), blk_of.end(), -1);
|
||||
std::fill(cell_grp.begin(), cell_grp.end(), -1);
|
||||
std::fill(grp_head.begin(), grp_head.end(), -1);
|
||||
|
||||
grp_next .clear();
|
||||
grp_first.clear();
|
||||
grp_slot0.clear();
|
||||
grp_slots.clear();
|
||||
grp_bid .clear();
|
||||
|
||||
oor = false;
|
||||
dup = false;
|
||||
|
||||
for (int64_t j = 0; j < n_kv; ++j) {
|
||||
if (cells.is_empty(j)) {
|
||||
continue;
|
||||
}
|
||||
|
||||
const int64_t idx = ranked ? rank[j] : cells.pos_get(j);
|
||||
const int64_t pb = idx/r;
|
||||
|
||||
if (pb >= n_blocks) {
|
||||
oor = true;
|
||||
continue;
|
||||
}
|
||||
|
||||
int32_t g = -1;
|
||||
|
||||
for (int32_t c = grp_head[pb]; c >= 0; c = grp_next[c]) {
|
||||
if (one_seq || cells.seq_get_all((uint32_t) grp_first[c]) == cells.seq_get_all((uint32_t) j)) {
|
||||
g = c;
|
||||
break;
|
||||
}
|
||||
}
|
||||
|
||||
if (g < 0) {
|
||||
g = (int32_t) grp_first.size();
|
||||
|
||||
grp_next .push_back(grp_head[pb]);
|
||||
grp_first.push_back((int32_t) j);
|
||||
grp_slot0.push_back(-1);
|
||||
grp_slots.push_back(0);
|
||||
grp_bid .push_back(-1);
|
||||
|
||||
grp_head[pb] = g;
|
||||
}
|
||||
|
||||
const uint64_t bit = uint64_t(1) << (idx%r);
|
||||
|
||||
dup |= (grp_slots[g] & bit) != 0;
|
||||
|
||||
cell_grp[j] = g;
|
||||
grp_slots[g] |= bit;
|
||||
|
||||
if (idx%r == 0) {
|
||||
grp_slot0[g] = (int32_t) j;
|
||||
}
|
||||
}
|
||||
};
|
||||
|
||||
group_cells();
|
||||
|
||||
// mrope repeats one position across an image, so rank cells instead of using the position
|
||||
if (dup && ubatch->is_pos_2d() && one_seq) {
|
||||
order.clear();
|
||||
order.reserve(n_kv);
|
||||
|
||||
for (int64_t j = 0; j < n_kv; ++j) {
|
||||
if (!cells.is_empty(j)) {
|
||||
order.push_back((int32_t) j);
|
||||
}
|
||||
}
|
||||
|
||||
// same total order the mrope causal mask uses: pos, then ext.y, then ext.x
|
||||
std::sort(order.begin(), order.end(), [&cells](int32_t a, int32_t b) {
|
||||
const llama_pos pa = cells.pos_get(a);
|
||||
const llama_pos pb = cells.pos_get(b);
|
||||
|
||||
if (pa != pb) {
|
||||
return pa < pb;
|
||||
}
|
||||
|
||||
const auto & ea = cells.ext_get(a);
|
||||
|
||||
return cells.ext_get(b).is_2d_gt(ea.x, ea.y);
|
||||
});
|
||||
|
||||
rank.assign(n_kv, -1);
|
||||
|
||||
for (int64_t k = 0; k < (int64_t) order.size(); ++k) {
|
||||
rank[order[k]] = (int32_t) k;
|
||||
}
|
||||
|
||||
ranked = true;
|
||||
|
||||
group_cells();
|
||||
}
|
||||
|
||||
GGML_ASSERT((!blk_bias || !oor) && "qsa: cell position runs past the cell window");
|
||||
|
||||
int32_t n_bid = 0;
|
||||
|
||||
for (int64_t pb = 0; pb < n_blocks; ++pb) {
|
||||
for (int32_t g = grp_head[pb]; g >= 0; g = grp_next[g]) {
|
||||
if (grp_slots[g] != slots_full) {
|
||||
continue;
|
||||
}
|
||||
|
||||
grp_bid[g] = n_bid++;
|
||||
|
||||
bid_idx .push_back((int32_t) (pb*r));
|
||||
bid_cell .push_back(grp_first[g]);
|
||||
bid_slot0.push_back(grp_slot0[g]);
|
||||
}
|
||||
}
|
||||
|
||||
GGML_ASSERT(n_bid <= n_blocks);
|
||||
|
||||
for (int32_t b = 0; b < n_bid; ++b) {
|
||||
int32_t sec_pos[4] = { bid_idx[b], bid_idx[b], bid_idx[b], bid_idx[b] };
|
||||
|
||||
if (ranked) {
|
||||
const int32_t c = bid_slot0[b];
|
||||
const llama_pos p = cells.pos_get(c);
|
||||
const auto & e = cells.ext_get(c);
|
||||
|
||||
sec_pos[0] = p;
|
||||
sec_pos[1] = e.y;
|
||||
sec_pos[2] = e.x;
|
||||
sec_pos[3] = p;
|
||||
}
|
||||
|
||||
for (int64_t sec = 0; sec < 4; ++sec) {
|
||||
dst_blk_pos[sec*(n_blocks*n_ns) + s*n_blocks + b] = sec_pos[sec];
|
||||
}
|
||||
}
|
||||
|
||||
// unpooled cells all point at one spare block. a spare block exists only when some
|
||||
// cell is unpooled: n_bid == n_blocks means every cell sits in a full block.
|
||||
const bool have_dead = n_bid < n_blocks;
|
||||
const int32_t dead_bid = have_dead ? n_bid : n_blocks - 1;
|
||||
|
||||
for (int64_t j = 0; j < n_kv; ++j) {
|
||||
const int32_t g = cell_grp[j];
|
||||
|
||||
blk_of[j] = g < 0 ? -1 : grp_bid[g];
|
||||
|
||||
if (blk_of[j] >= 0) {
|
||||
const int64_t idx = ranked ? rank[j] : cells.pos_get(j);
|
||||
|
||||
cur_blk_cells[blk_of[j]*r + (idx%r)] = (int32_t) j;
|
||||
}
|
||||
|
||||
cur_cell_blk[j] = blk_of[j] < 0 ? dead_bid : blk_of[j];
|
||||
}
|
||||
|
||||
for (int64_t ii = 0; ii < n_tps; ++ii) {
|
||||
const int64_t i = s*n_tps + ii;
|
||||
const llama_seq_id seq_id = ubatch->seq_id[i][0];
|
||||
|
||||
int64_t q = ubatch->pos[i];
|
||||
|
||||
if (ranked) {
|
||||
const llama_pos qt = ubatch->pos[i];
|
||||
const llama_pos qy = ubatch->pos[i + n_tokens];
|
||||
const llama_pos qx = ubatch->pos[i + n_tokens*2];
|
||||
|
||||
int64_t lo = 0;
|
||||
int64_t hi = (int64_t) order.size();
|
||||
|
||||
while (lo < hi) {
|
||||
const int64_t mid = (lo + hi)/2;
|
||||
const int32_t c = order[mid];
|
||||
const llama_pos pc = cells.pos_get(c);
|
||||
|
||||
if (pc < qt || (pc == qt && !cells.ext_get(c).is_2d_gt(qx, qy))) {
|
||||
lo = mid + 1;
|
||||
} else {
|
||||
hi = mid;
|
||||
}
|
||||
}
|
||||
|
||||
q = lo - 1;
|
||||
}
|
||||
|
||||
// the tail is an incomplete block and is always visible, as in the reference
|
||||
const int64_t tail_start = (q + 1)/r*r;
|
||||
|
||||
if (blk_bias) {
|
||||
// a block sits wholly inside or outside the tail, so one value covers it
|
||||
// the caller adds the attention mask, which drops empty, foreign and, when causal, future cells
|
||||
float * cur_blk_bias = dst_bias + i*n_blocks;
|
||||
|
||||
for (int64_t b = 0; b < n_blocks; ++b) {
|
||||
if (b >= n_bid || !cells.seq_has((uint32_t) bid_cell[b], seq_id)) {
|
||||
cur_blk_bias[b] = -INFINITY;
|
||||
continue;
|
||||
}
|
||||
|
||||
// finite, so it can never meet a -inf and produce a nan
|
||||
cur_blk_bias[b] = (causal_attn && bid_idx[b] >= tail_start) ? 1e9f : 0.0f;
|
||||
}
|
||||
|
||||
// the spare block holds the unpooled cells, which are the incomplete tail, so
|
||||
// it gets the tail value. it must stay finite: a sequence with fewer than
|
||||
// `ratio` cells owns no full block, and a row of -inf only gives a nan.
|
||||
if (have_dead) {
|
||||
cur_blk_bias[dead_bid] = 1e9f;
|
||||
}
|
||||
|
||||
continue;
|
||||
}
|
||||
|
||||
float * cur_bias = dst_bias + i*n_kv;
|
||||
|
||||
for (int64_t j = 0; j < n_kv; ++j) {
|
||||
float v = -INFINITY;
|
||||
|
||||
if (!cells.is_empty(j) && cells.seq_has(j, seq_id)) {
|
||||
const int64_t idx = ranked ? rank[j] : cells.pos_get(j);
|
||||
|
||||
if (!causal_attn) {
|
||||
// every visible block competes on score and the unpooled cells are always selected
|
||||
v = blk_of[j] < 0 ? 1e9f : 0.0f;
|
||||
} else if (idx <= q) {
|
||||
// finite, so it can never meet a -inf and produce a nan
|
||||
v = idx >= tail_start ? 1e9f : (blk_of[j] < 0 ? -INFINITY : 0.0f);
|
||||
}
|
||||
}
|
||||
|
||||
cur_bias[j] = v;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
//
|
||||
// llama_memory_hybrid_idx_context
|
||||
//
|
||||
@@ -697,6 +376,7 @@ struct llama_memory_hybrid_idx_context::kpool_state {
|
||||
|
||||
uint32_t n_pool_real = 0;
|
||||
uint32_t n_new = 0;
|
||||
uint32_t n_new_g = 1; // graph size of the new pool list, stable across decode steps
|
||||
bool cache_safe = true;
|
||||
};
|
||||
|
||||
@@ -707,6 +387,13 @@ uint32_t kpool_pad(uint32_t n_pool) {
|
||||
return std::max<uint32_t>(64u, GGML_PAD(n_pool + 1, 64u));
|
||||
}
|
||||
|
||||
// Rank of (pos, cell) in a sequence's cells sorted by position then cell, or -1 when absent.
|
||||
// In order mode the rank alone places a token: cells sharing a position (M-RoPE images) have distinct ranks.
|
||||
int64_t kpool_rank(const std::vector<std::pair<llama_pos, uint32_t>> & cells, llama_pos pos, uint32_t cell) {
|
||||
auto it = std::lower_bound(cells.begin(), cells.end(), std::make_pair(pos, cell));
|
||||
return it != cells.end() && it->second == cell && it->first == pos ? it - cells.begin() : -1;
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
llama_memory_hybrid_idx::~llama_memory_hybrid_idx() = default;
|
||||
@@ -787,24 +474,31 @@ const llama_memory_hybrid_idx::kpool_layout & llama_memory_hybrid_idx::kpool_lay
|
||||
|
||||
// Pools start at the first valid token
|
||||
size_t j = sq.j_next;
|
||||
while (j + kpool <= sq.cells.size()) {
|
||||
const llama_pos p0 = sq.cells[j].first;
|
||||
if ((p0 - sq.pos_min) % (llama_pos) kpool != 0) {
|
||||
++j;
|
||||
continue;
|
||||
}
|
||||
bool ok = true;
|
||||
for (uint32_t k = 1; k < kpool; ++k) {
|
||||
if (sq.cells[j + k].first != p0 + (llama_pos) k) {
|
||||
ok = false;
|
||||
break;
|
||||
}
|
||||
}
|
||||
if (ok) {
|
||||
if (hparams_idx.indexer_kpool_by_order) {
|
||||
// consecutive cells in sequence order, whatever their positions
|
||||
for (; j + kpool <= sq.cells.size(); j += kpool) {
|
||||
sq.pools.push_back((uint32_t) j);
|
||||
j += kpool;
|
||||
} else {
|
||||
++j;
|
||||
}
|
||||
} else {
|
||||
while (j + kpool <= sq.cells.size()) {
|
||||
const llama_pos p0 = sq.cells[j].first;
|
||||
if ((p0 - sq.pos_min) % (llama_pos) kpool != 0) {
|
||||
++j;
|
||||
continue;
|
||||
}
|
||||
bool ok = true;
|
||||
for (uint32_t k = 1; k < kpool; ++k) {
|
||||
if (sq.cells[j + k].first != p0 + (llama_pos) k) {
|
||||
ok = false;
|
||||
break;
|
||||
}
|
||||
}
|
||||
if (ok) {
|
||||
sq.pools.push_back((uint32_t) j);
|
||||
j += kpool;
|
||||
} else {
|
||||
++j;
|
||||
}
|
||||
}
|
||||
}
|
||||
sq.j_next = j;
|
||||
@@ -835,7 +529,8 @@ llama_memory_hybrid_idx_context::llama_memory_hybrid_idx_context(llama_memory_hy
|
||||
const uint64_t n_pool_max = uint64_t(idx->get_size() / mem->get_kpool()) * idx->get_n_seq_max();
|
||||
GGML_ASSERT(n_pool_max <= UINT32_MAX - 64);
|
||||
st.n_pool_real = std::max(st.n_pool_real, uint32_t(n_pool_max));
|
||||
st.n_new = st.n_pool_real;
|
||||
st.n_new = st.n_pool_real;
|
||||
st.n_new_g = std::max(st.n_new, 1u);
|
||||
kpool_st = std::make_unique<kpool_state>(std::move(st));
|
||||
i_kpool = 0;
|
||||
}
|
||||
@@ -860,6 +555,7 @@ llama_memory_hybrid_idx_context::llama_memory_hybrid_idx_context(
|
||||
llama_memory_hybrid_context(mem, std::move(sinfos_attn), ubatches),
|
||||
mem(mem),
|
||||
ns_ubatch(llama_memory_hybrid_idx_ns(sinfos_idx)),
|
||||
sinfos_kpool(mem->get_mem_idx() != nullptr && mem->get_kpool() > 0 && mem->get_kpool_by_order() ? sinfos_idx : slot_info_vec_t()),
|
||||
ctx_idx(mem->get_mem_idx() == nullptr ? nullptr :
|
||||
new llama_kv_cache_context(mem->get_mem_idx(), std::move(sinfos_idx), ubatches)) {
|
||||
// Sequence edits force the touched positions to re-pool.
|
||||
@@ -918,29 +614,17 @@ uint32_t llama_memory_hybrid_idx_context::get_n_stream() const {
|
||||
return ns_ubatch[i_cur];
|
||||
}
|
||||
|
||||
void llama_memory_hybrid_idx_context::set_input_qsa(
|
||||
ggml_tensor * cell_blk,
|
||||
ggml_tensor * blk_cells,
|
||||
ggml_tensor * blk_pos,
|
||||
ggml_tensor * bias,
|
||||
const llama_ubatch * ubatch,
|
||||
uint32_t ratio,
|
||||
bool blk_bias,
|
||||
bool causal_attn) const {
|
||||
GGML_ASSERT(mem != nullptr);
|
||||
|
||||
mem->set_input_qsa(cell_blk, blk_cells, blk_pos, bias, ubatch, ratio, blk_bias, causal_attn);
|
||||
}
|
||||
|
||||
llama_memory_hybrid_idx_context::kpool_access::kpool_access(ggml_context * ctx, ggml_tensor * k, int64_t n_embd) : ctx(ctx) {
|
||||
GGML_ASSERT(k->ne[0] == 3*n_embd);
|
||||
// rows are the per-token part (glm5-next: key | gate, qwen4exp: key), then the pooled key
|
||||
const int64_t n_tok = k->ne[0] - n_embd;
|
||||
GGML_ASSERT(n_tok > 0 && n_tok % n_embd == 0);
|
||||
|
||||
const int64_t n_cells = k->ne[1]*k->ne[2];
|
||||
|
||||
// Pool indices can refer to other streams. Revisit these full-storage views if that changes:
|
||||
// https://github.com/ggml-org/llama.cpp/pull/27773#discussion_r4130905603
|
||||
key_gate = ggml_view_2d(ctx, k, 2*n_embd, n_cells, k->nb[1], 0);
|
||||
pooled = ggml_view_2d(ctx, k, n_embd, n_cells, k->nb[1], ggml_row_size(k->type, 2*n_embd));
|
||||
key_gate = ggml_view_2d(ctx, k, n_tok, n_cells, k->nb[1], 0);
|
||||
pooled = ggml_view_2d(ctx, k, n_embd, n_cells, k->nb[1], ggml_row_size(k->type, n_tok));
|
||||
}
|
||||
|
||||
ggml_tensor * llama_memory_hybrid_idx_context::kpool_access::gather_key_gate(ggml_tensor * idxs) const {
|
||||
@@ -972,7 +656,7 @@ ggml_tensor * llama_memory_hybrid_idx_context::gather_mla_rows(
|
||||
return ggml_get_rows(ctx, rows, ggml_reshape_1d(ctx, idxs, n_rows));
|
||||
}
|
||||
|
||||
// k-pool DSA indexer (glm5-next)
|
||||
// k-pool DSA indexer (glm5-next, qwen4exp QSA)
|
||||
|
||||
// Sizes only, used by the full cache context so get_n_kpool() works during graph reserve.
|
||||
llama_memory_hybrid_idx_context::kpool_state llama_memory_hybrid_idx_context::kpool_build_sizes() const {
|
||||
@@ -1034,7 +718,7 @@ void llama_memory_hybrid_idx_context::kpool_build_state(const llama_ubatch & uba
|
||||
}
|
||||
|
||||
auto first = std::lower_bound(sq.pools.begin(), sq.pools.end(), stale_from,
|
||||
[&](uint32_t j, llama_pos p) { return sq.cells[j].first + (llama_pos) kpool <= p; });
|
||||
[&](uint32_t j, llama_pos p) { return sq.cells[j + kpool - 1].first < p; });
|
||||
for (auto it = first; it != sq.pools.end(); ++it) {
|
||||
mark(pool_start[s] + (uint32_t) (it - sq.pools.begin()));
|
||||
}
|
||||
@@ -1043,26 +727,45 @@ void llama_memory_hybrid_idx_context::kpool_build_state(const llama_ubatch & uba
|
||||
|
||||
if (!st.cache_safe) {
|
||||
std::fill(st.is_new.begin(), st.is_new.end(), st.generation);
|
||||
st.n_new = st.n_pool_real;
|
||||
st.n_new = st.n_pool_real;
|
||||
st.n_new_g = std::max(st.n_new, 1u);
|
||||
return;
|
||||
}
|
||||
|
||||
// in order mode a token's cell gives its rank, and the rank its pool: positions cannot, as an image shares one
|
||||
const bool by_order = mem->get_kpool_by_order();
|
||||
const auto * sinfo = by_order ? &sinfos_kpool[i_cur] : nullptr;
|
||||
const uint32_t n_tps = by_order ? (uint32_t) sinfo->size() : 0;
|
||||
|
||||
for (uint32_t i = 0; i < ubatch.n_tokens; ++i) {
|
||||
const llama_pos p = ubatch.pos[i];
|
||||
for (int32_t k = 0; k < ubatch.n_seq_id[i]; ++k) {
|
||||
const llama_seq_id s = ubatch.seq_id[i][k];
|
||||
const auto & sq = lay.seqs[s];
|
||||
if (by_order) {
|
||||
const int64_t r = kpool_rank(sq.cells, p, sinfo->idxs[i / n_tps][i % n_tps]);
|
||||
GGML_ASSERT(r >= 0);
|
||||
if ((size_t) r / kpool < sq.pools.size()) {
|
||||
mark(pool_start[s] + (uint32_t) (r / kpool));
|
||||
}
|
||||
continue;
|
||||
}
|
||||
auto it = std::upper_bound(sq.pools.begin(), sq.pools.end(), p,
|
||||
[&](llama_pos pos, uint32_t j) { return pos < sq.cells[j].first; });
|
||||
if (it == sq.pools.begin()) {
|
||||
continue;
|
||||
}
|
||||
--it;
|
||||
if (p < sq.cells[*it].first + (llama_pos) kpool) {
|
||||
if (p <= sq.cells[*it + kpool - 1].first) {
|
||||
mark(pool_start[s] + (uint32_t) (it - sq.pools.begin()));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// a ubatch touches at most t_s/kpool + 1 pools of a sequence with t_s tokens: pad to that bound so the
|
||||
// graph keeps its shape as the count moves, e.g. between 0 and n_seq while several sequences decode
|
||||
const uint32_t bound = ubatch.n_tokens/kpool + ubatch.n_seqs_unq;
|
||||
st.n_new_g = std::max({st.n_new, 1u, std::min(bound, kpool_pad(st.n_pool_real) - 1)});
|
||||
}
|
||||
|
||||
const llama_memory_hybrid_idx_context::kpool_state & llama_memory_hybrid_idx_context::kpool_cur() const {
|
||||
@@ -1076,7 +779,7 @@ uint32_t llama_memory_hybrid_idx_context::get_n_kpool() const {
|
||||
}
|
||||
|
||||
uint32_t llama_memory_hybrid_idx_context::get_n_kpool_new() const {
|
||||
return kpool_cur().n_new;
|
||||
return kpool_cur().n_new_g;
|
||||
}
|
||||
|
||||
bool llama_memory_hybrid_idx_context::get_kpool_cache_safe() const {
|
||||
@@ -1085,7 +788,7 @@ bool llama_memory_hybrid_idx_context::get_kpool_cache_safe() const {
|
||||
|
||||
void llama_memory_hybrid_idx_context::set_input_kpool(ggml_tensor * pool_cells, ggml_tensor * pool_idxs, ggml_tensor * pool_mask, ggml_tensor * tail_idxs,
|
||||
ggml_tensor * gather_mask, bool gather, ggml_tensor * new_pool_idxs, ggml_tensor * new_pool_rep,
|
||||
const llama_ubatch * ubatch) const {
|
||||
const llama_ubatch * ubatch, ggml_tensor * new_pool_pos) const {
|
||||
GGML_ASSERT(mem != nullptr && mem->get_mem_idx() != nullptr);
|
||||
GGML_ASSERT(ggml_backend_buffer_is_host(pool_cells->buffer));
|
||||
GGML_ASSERT(ggml_backend_buffer_is_host(pool_idxs->buffer));
|
||||
@@ -1101,8 +804,10 @@ void llama_memory_hybrid_idx_context::set_input_kpool(ggml_tensor * pool_cells,
|
||||
const uint32_t n_tokens = ubatch->n_tokens;
|
||||
const uint32_t n_pool = (uint32_t) pool_cells->ne[0];
|
||||
const uint32_t n_new = st.n_new;
|
||||
// the graph always pools at least one entry, see build_inp_kpool
|
||||
const uint32_t n_new_g = std::max(n_new, 1u);
|
||||
// the graph always pools at least one entry, padded to a stable bound, see kpool_build_state
|
||||
const uint32_t n_new_g = st.n_new_g;
|
||||
|
||||
const bool by_order = mem->get_kpool_by_order();
|
||||
|
||||
GGML_ASSERT(n_pool == kpool_pad(st.n_pool_real));
|
||||
GGML_ASSERT(st.is_new.size() == st.n_pool_real);
|
||||
@@ -1116,6 +821,10 @@ void llama_memory_hybrid_idx_context::set_input_kpool(ggml_tensor * pool_cells,
|
||||
GGML_ASSERT(ggml_backend_buffer_is_host(new_pool_rep->buffer));
|
||||
GGML_ASSERT(new_pool_rep->ne[0] == (int64_t) n_new_g);
|
||||
}
|
||||
if (new_pool_pos != nullptr) {
|
||||
GGML_ASSERT(ggml_backend_buffer_is_host(new_pool_pos->buffer));
|
||||
GGML_ASSERT(new_pool_pos->ne[0] == 4*(int64_t) n_new_g);
|
||||
}
|
||||
|
||||
const uint32_t kv_size = mem->get_mem_idx()->get_size();
|
||||
const uint32_t n_stream_kv = mem->get_mem_idx()->get_n_stream();
|
||||
@@ -1142,6 +851,19 @@ void llama_memory_hybrid_idx_context::set_input_kpool(ggml_tensor * pool_cells,
|
||||
dummy_cell = gcell(sq, it->second);
|
||||
}
|
||||
|
||||
// in order mode a token sees the pools and the tail up to its own rank in the sequence, which its cell pins down
|
||||
std::vector<int64_t> rank;
|
||||
if (by_order) {
|
||||
const auto & sinfo = sinfos_kpool[i_cur];
|
||||
const uint32_t n_tps = (uint32_t) sinfo.size();
|
||||
|
||||
rank.resize(n_tokens);
|
||||
for (uint32_t i = 0; i < n_tokens; ++i) {
|
||||
rank[i] = kpool_rank(lay.seqs[ubatch->seq_id[i][0]].cells, ubatch->pos[i], sinfo.idxs[i / n_tps][i % n_tps]);
|
||||
GGML_ASSERT(rank[i] >= 0);
|
||||
}
|
||||
}
|
||||
|
||||
// Gather maps padding to a real cell and masks it separately.
|
||||
const int32_t sentinel = gather ? (int32_t) dummy_cell : (int32_t) n_kv;
|
||||
|
||||
@@ -1168,6 +890,11 @@ void llama_memory_hybrid_idx_context::set_input_kpool(ggml_tensor * pool_cells,
|
||||
int32_t * pidx = (int32_t *) pool_idxs->data;
|
||||
int32_t * nidx = (int32_t *) new_pool_idxs->data;
|
||||
int64_t * nrep = new_pool_rep != nullptr ? (int64_t *) new_pool_rep->data : nullptr;
|
||||
int32_t * npos = new_pool_pos != nullptr ? (int32_t *) new_pool_pos->data : nullptr;
|
||||
|
||||
if (npos != nullptr) {
|
||||
std::fill(npos, npos + 4*n_new_g, 0);
|
||||
}
|
||||
|
||||
uint32_t i_new = 0;
|
||||
for (llama_seq_id s = 0; s < LLAMA_MAX_SEQ; ++s) {
|
||||
@@ -1198,6 +925,15 @@ void llama_memory_hybrid_idx_context::set_input_kpool(ggml_tensor * pool_cells,
|
||||
if (nrep != nullptr) {
|
||||
nrep[i_new] = gcell(sq, rep);
|
||||
}
|
||||
if (npos != nullptr) {
|
||||
// a pooled key is rotated to the M-RoPE position of its first member
|
||||
const uint32_t c = sq.cells[j].second;
|
||||
const auto & e = mem->get_mem_idx()->get_cells(s).ext_get(c);
|
||||
npos[0*n_new_g + i_new] = sq.cells[j].first;
|
||||
npos[1*n_new_g + i_new] = e.y;
|
||||
npos[2*n_new_g + i_new] = e.x;
|
||||
npos[3*n_new_g + i_new] = sq.cells[j].first;
|
||||
}
|
||||
++i_new;
|
||||
}
|
||||
|
||||
@@ -1206,14 +942,25 @@ void llama_memory_hybrid_idx_context::set_input_kpool(ggml_tensor * pool_cells,
|
||||
}
|
||||
GGML_ASSERT(i_new == n_new);
|
||||
|
||||
// A ubatch that completes no pool re-pools the cell of its first token. That cell cannot belong to
|
||||
// a complete pool here, else the pool would be marked new, so the write never touches a cached key.
|
||||
if (n_new == 0) {
|
||||
// Padded entries re-pool a cell whose pooled slot is never read. With no new pool that is the cell of the
|
||||
// first token: it cannot belong to a complete pool, else the pool would be marked new. Otherwise it is the
|
||||
// first member of a pool, which is never a pool's rep.
|
||||
int64_t pad_cell = dummy_cell;
|
||||
if (n_new > 0) {
|
||||
for (llama_seq_id s = 0; s < LLAMA_MAX_SEQ; ++s) {
|
||||
const auto & sq = lay.seqs[s];
|
||||
if (!sq.pools.empty()) {
|
||||
pad_cell = gcell(sq, sq.cells[sq.pools[0]].second);
|
||||
break;
|
||||
}
|
||||
}
|
||||
}
|
||||
for (uint32_t i = n_new; i < n_new_g; ++i) {
|
||||
for (uint32_t k = 0; k < kpool; ++k) {
|
||||
nidx[k] = (int32_t) dummy_cell;
|
||||
nidx[(size_t) i*kpool + k] = (int32_t) pad_cell;
|
||||
}
|
||||
if (nrep != nullptr) {
|
||||
nrep[0] = dummy_cell;
|
||||
nrep[i] = pad_cell;
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1240,7 +987,8 @@ void llama_memory_hybrid_idx_context::set_input_kpool(ggml_tensor * pool_cells,
|
||||
|
||||
const uint32_t p0 = seq_pool_start[s];
|
||||
const uint32_t p1 = p0 + (uint32_t) lay.seqs[s].pools.size();
|
||||
const uint32_t nv = (uint32_t) (std::upper_bound(pool_end.begin() + p0, pool_end.begin() + p1, p) - (pool_end.begin() + p0));
|
||||
const uint32_t nv = by_order ? std::min(p1 - p0, (uint32_t) ((rank[i] + 1)/kpool)) :
|
||||
(uint32_t) (std::upper_bound(pool_end.begin() + p0, pool_end.begin() + p1, p) - (pool_end.begin() + p0));
|
||||
std::fill(row + p0, row + p0 + nv, keep);
|
||||
|
||||
// Finite visible pools occupy the first min(nv, n_top) ranked slots.
|
||||
@@ -1264,12 +1012,18 @@ void llama_memory_hybrid_idx_context::set_input_kpool(ggml_tensor * pool_cells,
|
||||
const llama_pos p = ubatch->pos[i];
|
||||
const auto & sq = lay.seqs[s];
|
||||
|
||||
const uint32_t n_tail = (uint32_t) ((p - sq.pos_min + 1) % (llama_pos) kpool);
|
||||
const uint32_t n_tail = by_order ?
|
||||
(uint32_t) ((rank[i] + 1) % kpool) :
|
||||
(uint32_t) ((p - sq.pos_min + 1) % (llama_pos) kpool);
|
||||
|
||||
for (uint32_t k = 0; k < kpool - 1; ++k) {
|
||||
int32_t cell = sentinel;
|
||||
bool real = false;
|
||||
if (k < n_tail) {
|
||||
if (k < n_tail && by_order) {
|
||||
const uint32_t c = sq.cells[rank[i] - k].second;
|
||||
cell = (int32_t) (gather ? gcell(sq, c) : (int64_t) c);
|
||||
real = true;
|
||||
} else if (k < n_tail) {
|
||||
const llama_pos pt = p - (llama_pos) k;
|
||||
auto it = std::lower_bound(sq.cells.begin(), sq.cells.end(), std::make_pair(pt, 0u));
|
||||
if (it != sq.cells.end() && it->first == pt) {
|
||||
|
||||
@@ -80,23 +80,13 @@ public:
|
||||
|
||||
llama_kv_cache * get_mem_idx() const; // nullptr when the model carries no indexer
|
||||
|
||||
// block-compressed sparse attention (qwen4exp QSA) over the cells of the indexer cache.
|
||||
// Blocks cut the position line, not the cell array, so no caller assumes a contiguous layout:
|
||||
// cell_blk I32 [n_kv, ns] block each cell belongs to
|
||||
// blk_cells I32 [ratio*n_blocks, ns] cells making up each block
|
||||
// blk_pos I32 [4*n_blocks*ns] mrope position rows of each block's first token
|
||||
// bias F32 [n_kv, n_tokens/ns, ns] -inf where invisible, large where always visible
|
||||
// blk_bias asks for the bias per block instead: [n_blocks, n_tokens/ns, ns]
|
||||
// the caller then adds the attention mask, the only part of the bias that varies within a block
|
||||
// causal_attn selects the rule: causal forces the query's own block on, non-causal lets every visible block compete on score
|
||||
void set_input_qsa(ggml_tensor * cell_blk, ggml_tensor * blk_cells, ggml_tensor * blk_pos,
|
||||
ggml_tensor * bias, const llama_ubatch * ubatch, uint32_t ratio,
|
||||
bool blk_bias, bool causal_attn) const;
|
||||
|
||||
// The model's indexer pool size.
|
||||
uint32_t get_kpool() const { return hparams_idx.indexer_kpool; }
|
||||
|
||||
// Which cells of a sequence make up which pool of kpool consecutive positions.
|
||||
// Whether pools are kpool consecutive cells in sequence order (qwen4exp) instead of kpool consecutive positions.
|
||||
bool get_kpool_by_order() const { return hparams_idx.indexer_kpool_by_order; }
|
||||
|
||||
// Which cells of a sequence make up which pool of kpool consecutive positions (or cells, in order mode).
|
||||
// It is kept here because it outlives the batch: pools are fixed by the positions relative to the
|
||||
// sequence's first one, so a ubatch only ever appends to it. Sequence edits drop it, see mem_idx_stale.
|
||||
struct kpool_layout;
|
||||
@@ -203,18 +193,16 @@ public:
|
||||
// streams in the current slot info, the `ns` of get_k/get_v; 1 if unified
|
||||
uint32_t get_n_stream() const;
|
||||
|
||||
// glm5-next, complete pools of kpool consecutive positions per sequence, scored as whole pools.
|
||||
// glm5-next and qwen4exp, complete pools of kpool cells per sequence, scored as whole pools.
|
||||
uint32_t get_n_kpool () const; // Padded pool count, where the last pool is always unused.
|
||||
uint32_t get_n_kpool_new() const; // Exact count of pools completed by the current ubatch.
|
||||
uint32_t get_n_kpool_new() const; // Pools to re-pool this ubatch, padded to a stable bound, never below 1.
|
||||
bool get_kpool_cache_safe() const;
|
||||
kpool_access get_kpool_access(ggml_context * ctx, int32_t il, int64_t n_embd) const;
|
||||
ggml_tensor * gather_mla_rows(ggml_context * ctx, ggml_tensor * idxs, int64_t n_rows, int64_t n_embd, int32_t il) const;
|
||||
// new_pool_pos (I32 [4*n_new]): M-RoPE position of each new pool's first member, for pooled keys rotated at pooling time
|
||||
void set_input_kpool(ggml_tensor * pool_cells, ggml_tensor * pool_idxs, ggml_tensor * pool_mask, ggml_tensor * tail_idxs,
|
||||
ggml_tensor * gather_mask, bool gather, ggml_tensor * new_pool_idxs, ggml_tensor * new_pool_rep,
|
||||
const llama_ubatch * ubatch) const;
|
||||
void set_input_qsa(ggml_tensor * cell_blk, ggml_tensor * blk_cells, ggml_tensor * blk_pos,
|
||||
ggml_tensor * bias, const llama_ubatch * ubatch, uint32_t ratio,
|
||||
bool blk_bias, bool causal_attn) const;
|
||||
const llama_ubatch * ubatch, ggml_tensor * new_pool_pos = nullptr) const;
|
||||
|
||||
private:
|
||||
llama_memory_hybrid_idx * mem = nullptr;
|
||||
@@ -223,6 +211,10 @@ private:
|
||||
// declared first, so it is initialised while sinfos_idx is still intact
|
||||
const std::vector<uint32_t> ns_ubatch;
|
||||
|
||||
// the indexer cells of each ubatch, kept for pools in cache order (qwen4exp): token s*n + i of ubatch u
|
||||
// sits in cell idxs[s][i] of stream strm[s] of sinfos_kpool[u], and several cells can share a position
|
||||
const slot_info_vec_t sinfos_kpool;
|
||||
|
||||
// null unless the model has an indexer
|
||||
const llama_memory_context_ptr ctx_idx;
|
||||
|
||||
|
||||
+10
-8
@@ -2388,7 +2388,7 @@ struct llama_model_qwen35 : public llama_model_base {
|
||||
struct llama_model_qwen4exp : public llama_model_base {
|
||||
llama_model_qwen4exp(const struct llama_model_params & params) : llama_model_base(params) {}
|
||||
|
||||
class llm_graph_input_qsa;
|
||||
class llm_graph_input_kpool;
|
||||
|
||||
void load_arch_hparams(llama_model_loader & ml) override;
|
||||
void load_arch_tensors(llama_model_loader & ml) override;
|
||||
@@ -2415,28 +2415,30 @@ struct llama_model_qwen4exp : public llama_model_base {
|
||||
ggml_tensor * build_layer_attn(
|
||||
llm_graph_input_attn_kv * inp_attn,
|
||||
const llama_memory_hybrid_idx_context * mctx_hyb,
|
||||
llm_graph_input_kpool * inp_kpool,
|
||||
ggml_tensor * cur,
|
||||
ggml_tensor * inp_pos,
|
||||
int * sections,
|
||||
int il);
|
||||
|
||||
// dense self-attention restricted to the cells that top_k names
|
||||
// dense self-attention over the cells the QSA mask keeps
|
||||
ggml_tensor * build_attn_qsa(
|
||||
llm_graph_input_attn_kv * inp,
|
||||
ggml_tensor * q_cur,
|
||||
ggml_tensor * k_cur,
|
||||
ggml_tensor * v_cur,
|
||||
ggml_tensor * top_k,
|
||||
ggml_tensor * sel,
|
||||
int64_t n_sel,
|
||||
float kq_scale,
|
||||
int il);
|
||||
|
||||
// the QSA cache layout inputs do not depend on the layer, only on its compress ratio,
|
||||
// so the layers sharing a ratio share one input set
|
||||
std::map<uint32_t, llm_graph_input_qsa *> qsa_inps;
|
||||
// the QSA layers share one set of k-pool inputs, see llama_memory_hybrid_idx
|
||||
llm_graph_input_kpool * build_inp_kpool(const llama_memory_hybrid_idx_context * mctx_hyb);
|
||||
|
||||
// QSA: token indices this layer's queries may attend to, or nullptr for dense
|
||||
ggml_tensor * build_qsa_top_k(
|
||||
// QSA: the additive mask [n_kv, n_tokens] of the top blocks and the tail, kq_mask included
|
||||
ggml_tensor * build_qsa_sel(
|
||||
const llama_memory_hybrid_idx_context * mctx_hyb,
|
||||
llm_graph_input_kpool * inp_kpool,
|
||||
ggml_tensor * cur,
|
||||
ggml_tensor * inp_pos,
|
||||
ggml_tensor * kq_mask,
|
||||
|
||||
+201
-181
@@ -64,6 +64,27 @@ void llama_model_qwen4exp::load_arch_hparams(llama_model_loader & ml) {
|
||||
qwen4exp_require_nonzero(ml, LLM_KV_ATTENTION_INDEXER_TOP_K, hparams.indexer_top_k);
|
||||
ml.get_key_or_arr(LLM_KV_ATTENTION_COMPRESS_RATIOS, hparams.dsv4_compress_ratios, hparams.n_layer_all, false);
|
||||
|
||||
// QSA pools the indexer keys of blocks of compress_ratio cells, one block size for the whole model
|
||||
hparams.indexer_kpool = 0;
|
||||
for (uint32_t il = 0; il < hparams.n_layer(); ++il) {
|
||||
const uint32_t r = hparams.dsv4_compress_ratios[il];
|
||||
if (r == 0) {
|
||||
continue;
|
||||
}
|
||||
if (hparams.indexer_kpool != 0 && r != hparams.indexer_kpool) {
|
||||
throw std::runtime_error(format("QSA layers must share one compress ratio, got %u and %u", hparams.indexer_kpool, r));
|
||||
}
|
||||
hparams.indexer_kpool = r;
|
||||
}
|
||||
if (hparams.indexer_kpool == 1 || (hparams.indexer_kpool > 0 && hparams.indexer_top_k % hparams.indexer_kpool != 0)) {
|
||||
throw std::runtime_error(format("QSA needs a compress ratio above 1 that divides the budget, got %u and %u",
|
||||
hparams.indexer_kpool, hparams.indexer_top_k));
|
||||
}
|
||||
// the reference groups the visible tokens in cache order and always keeps the tail
|
||||
hparams.indexer_kpool_row = 2; // raw key | pooled key
|
||||
hparams.indexer_kpool_by_order = true;
|
||||
hparams.indexer_kpool_select_tail = true;
|
||||
|
||||
// PLE n-gram hash embeddings; if the key group is absent every field stays zero
|
||||
hparams.is_ple_impl.reset();
|
||||
hparams.ple_n_heads = 0;
|
||||
@@ -378,6 +399,13 @@ llama_model_qwen4exp::graph::graph(const llama_model & model, const llm_graph_pa
|
||||
"the indexer cache must track the attention cache cell for cell");
|
||||
}
|
||||
|
||||
// the QSA layers share one set of k-pool inputs
|
||||
// the CUDA lightning indexer takes 32 or 64 heads, QSA has a few, so it scores with plain ops
|
||||
llm_graph_input_kpool * inp_kpool = nullptr;
|
||||
if (mctx_idx && hparams.indexer_kpool > 0) {
|
||||
inp_kpool = build_inp_kpool(mctx_hyb);
|
||||
}
|
||||
|
||||
ggml_tensor * inp_pos = build_inp_pos();
|
||||
ggml_tensor * inp_out_ids = build_inp_out_ids();
|
||||
|
||||
@@ -416,7 +444,7 @@ llama_model_qwen4exp::graph::graph(const llama_model & model, const llm_graph_pa
|
||||
if (hparams.is_recr(il)) {
|
||||
cur = build_layer_attn_linear(inp->get_recr(), cur, il);
|
||||
} else {
|
||||
cur = build_layer_attn(inp->get_attn(), mctx_hyb, cur, inp_pos, sections, il);
|
||||
cur = build_layer_attn(inp->get_attn(), mctx_hyb, inp_kpool, cur, inp_pos, sections, il);
|
||||
}
|
||||
|
||||
if (il == n_layer - 1 && inp_out_ids) {
|
||||
@@ -490,17 +518,16 @@ ggml_tensor * llama_model_qwen4exp::graph::build_norm_gated(
|
||||
return ggml_mul(ctx0, normalized, gated);
|
||||
}
|
||||
|
||||
// QSA attends to a budget of whole blocks of compress_ratio tokens, plus the incomplete tail
|
||||
// one mean-pooled indexer key scores each block; set_input resolves the cache layout
|
||||
class llama_model_qwen4exp::llm_graph_input_qsa : public llm_graph_input_i {
|
||||
// QSA k-pool inputs, shared by the QSA layers: blocks of compress_ratio cells in sequence order, see llama_memory_hybrid_idx
|
||||
class llama_model_qwen4exp::llm_graph_input_kpool : public llm_graph_input_i {
|
||||
public:
|
||||
llm_graph_input_qsa(const llama_memory_hybrid_idx_context * mctx, uint32_t ratio, bool blk_bias, bool causal_attn) :
|
||||
mctx(mctx), ratio(ratio), blk_bias(blk_bias), causal_attn(causal_attn) {}
|
||||
virtual ~llm_graph_input_qsa() = default;
|
||||
llm_graph_input_kpool(const llama_memory_hybrid_idx_context * mctx, uint32_t kpool) : mctx(mctx), kpool(kpool) {}
|
||||
virtual ~llm_graph_input_kpool() = default;
|
||||
|
||||
void set_input(const llama_ubatch * ubatch) override {
|
||||
mctx->get_idx()->set_input_k_idxs(k_idxs, ubatch);
|
||||
mctx->set_input_qsa(cell_blk, blk_cells, blk_pos, bias, ubatch, ratio, blk_bias, causal_attn);
|
||||
mctx->set_input_kpool(pool_cells, pool_idxs, pool_mask, tail_idxs, nullptr, false, new_pool_idxs, new_pool_rep,
|
||||
ubatch, new_pool_pos);
|
||||
}
|
||||
|
||||
bool can_reuse(const llm_graph_params & params) override {
|
||||
@@ -511,44 +538,85 @@ public:
|
||||
return false;
|
||||
}
|
||||
|
||||
const int64_t n_kv = idx->get_n_kv();
|
||||
const int64_t n_stream = mctx->get_n_stream();
|
||||
const int64_t n_blocks = (n_kv + ratio - 1)/ratio;
|
||||
|
||||
bool res = true;
|
||||
|
||||
res &= params.ubatch.n_tokens % n_stream == 0;
|
||||
|
||||
res &= k_idxs->ne[0] == params.ubatch.n_tokens;
|
||||
res &= cell_blk->ne[0] == n_kv;
|
||||
res &= cell_blk->ne[1] == n_stream;
|
||||
res &= blk_cells->ne[0] == (int64_t) ratio*n_blocks;
|
||||
res &= blk_pos->ne[0] == 4*n_blocks*n_stream;
|
||||
res &= bias->ne[0] == (blk_bias ? n_blocks : n_kv);
|
||||
res &= bias->ne[1] == params.ubatch.n_tokens/n_stream;
|
||||
res &= k_idxs->ne[0] == params.ubatch.n_tokens;
|
||||
res &= pool_cells->ne[0] == mctx->get_n_kpool();
|
||||
res &= pool_mask->ne[1] == params.ubatch.n_tokens;
|
||||
res &= tail_idxs->ne[1] == params.ubatch.n_tokens;
|
||||
// the scatter mask shape follows n_kv
|
||||
res &= n_kv == idx->get_n_kv();
|
||||
res &= n_new == mctx->get_n_kpool_new();
|
||||
res &= cache_safe == mctx->get_kpool_cache_safe();
|
||||
|
||||
return res;
|
||||
}
|
||||
|
||||
// per stream: a cell index names a different token in each stream
|
||||
ggml_tensor * k_idxs = nullptr; // I32 [n_tokens]
|
||||
ggml_tensor * cell_blk = nullptr; // I32 [n_kv, n_stream]
|
||||
ggml_tensor * blk_cells = nullptr; // I32 [ratio*n_blocks, n_stream]
|
||||
ggml_tensor * blk_pos = nullptr; // I32 [4*n_blocks*n_stream]
|
||||
ggml_tensor * bias = nullptr; // F32 [n_blocks or n_kv, n_tokens/n_stream, n_stream]
|
||||
ggml_tensor * k_idxs = nullptr; // I64 [n_tokens]
|
||||
ggml_tensor * pool_cells = nullptr; // I32 [n_pool] cell caching each block's pooled key
|
||||
ggml_tensor * pool_idxs = nullptr; // I32 [kpool, n_pool] member cells per block, n_kv sentinel for the padded blocks
|
||||
ggml_tensor * pool_mask = nullptr; // F32 [n_pool, n_tokens]
|
||||
ggml_tensor * tail_idxs = nullptr; // I32 [kpool - 1, n_tokens]
|
||||
ggml_tensor * new_pool_idxs = nullptr; // I32 [kpool, n_new] members of the blocks to re-pool this ubatch
|
||||
ggml_tensor * new_pool_rep = nullptr; // I64 [n_new] cell to write each new pooled key into
|
||||
ggml_tensor * new_pool_pos = nullptr; // I32 [4*n_new] M-RoPE position of each new block's first member
|
||||
|
||||
const llama_memory_hybrid_idx_context * mctx;
|
||||
const uint32_t ratio;
|
||||
|
||||
// the per-cell half of the bias is the attention mask, so only the per-block half is uploaded
|
||||
const bool blk_bias;
|
||||
|
||||
// this is fixed for the graph's lifetime, as causal_attn is part of the reuse key (llm_graph_params::allow_reuse)
|
||||
const bool causal_attn;
|
||||
const uint32_t kpool;
|
||||
uint32_t n_new = 0; // padded to a stable bound, never below 1
|
||||
uint32_t n_sel = 0;
|
||||
uint32_t n_kv = 0;
|
||||
bool cache_safe = true;
|
||||
};
|
||||
|
||||
ggml_tensor * llama_model_qwen4exp::graph::build_qsa_top_k(
|
||||
llama_model_qwen4exp::llm_graph_input_kpool * llama_model_qwen4exp::graph::build_inp_kpool(const llama_memory_hybrid_idx_context * mctx_hyb) {
|
||||
const auto * mctx_idx = mctx_hyb->get_idx();
|
||||
GGML_ASSERT(mctx_idx != nullptr);
|
||||
|
||||
const uint32_t kpool = hparams.indexer_kpool;
|
||||
const uint32_t n_pool = mctx_hyb->get_n_kpool();
|
||||
|
||||
auto inp = std::make_unique<llm_graph_input_kpool>(mctx_hyb, kpool);
|
||||
|
||||
inp->k_idxs = mctx_idx->build_input_k_idxs(ctx0, ubatch);
|
||||
inp->pool_cells = ggml_new_tensor_1d(ctx0, GGML_TYPE_I32, n_pool);
|
||||
inp->pool_idxs = ggml_new_tensor_2d(ctx0, GGML_TYPE_I32, kpool, n_pool);
|
||||
inp->pool_mask = ggml_new_tensor_2d(ctx0, GGML_TYPE_F32, n_pool, n_tokens);
|
||||
inp->tail_idxs = ggml_new_tensor_2d(ctx0, GGML_TYPE_I32, kpool - 1, n_tokens);
|
||||
ggml_set_input(inp->pool_cells);
|
||||
ggml_set_input(inp->pool_idxs);
|
||||
ggml_set_input(inp->pool_mask);
|
||||
ggml_set_input(inp->tail_idxs);
|
||||
|
||||
// set_input fills them all, so keep them allocated even when no op reads them
|
||||
ggml_build_forward_expand(gf, inp->pool_cells);
|
||||
ggml_build_forward_expand(gf, inp->pool_idxs);
|
||||
ggml_build_forward_expand(gf, inp->pool_mask);
|
||||
ggml_build_forward_expand(gf, inp->tail_idxs);
|
||||
|
||||
inp->n_kv = mctx_idx->get_n_kv();
|
||||
inp->n_new = mctx_hyb->get_n_kpool_new();
|
||||
inp->cache_safe = mctx_hyb->get_kpool_cache_safe();
|
||||
// the top blocks plus the tail
|
||||
inp->n_sel = kpool*std::min<uint32_t>(n_pool, hparams.indexer_top_k / kpool) + kpool - 1;
|
||||
|
||||
inp->new_pool_idxs = ggml_new_tensor_2d(ctx0, GGML_TYPE_I32, kpool, inp->n_new);
|
||||
ggml_set_input(inp->new_pool_idxs);
|
||||
if (inp->cache_safe) {
|
||||
inp->new_pool_rep = ggml_new_tensor_1d(ctx0, GGML_TYPE_I64, inp->n_new);
|
||||
ggml_set_input(inp->new_pool_rep);
|
||||
}
|
||||
inp->new_pool_pos = ggml_new_tensor_1d(ctx0, GGML_TYPE_I32, 4*inp->n_new);
|
||||
ggml_set_input(inp->new_pool_pos);
|
||||
|
||||
return (llm_graph_input_kpool *) res->add_input(std::move(inp));
|
||||
}
|
||||
|
||||
// QSA attends to the top blocks of compress_ratio cells plus the incomplete tail, like the glm5-next k-pool indexer
|
||||
// a block is scored by one pooled key: the mean of its raw indexer keys, normed and rotated to its first member
|
||||
ggml_tensor * llama_model_qwen4exp::graph::build_qsa_sel(
|
||||
const llama_memory_hybrid_idx_context * mctx_hyb,
|
||||
llm_graph_input_kpool * inp_kpool,
|
||||
ggml_tensor * cur,
|
||||
ggml_tensor * inp_pos,
|
||||
ggml_tensor * kq_mask,
|
||||
@@ -556,89 +624,57 @@ ggml_tensor * llama_model_qwen4exp::graph::build_qsa_top_k(
|
||||
int il) {
|
||||
const llama_kv_cache_context * mctx_idx = mctx_hyb->get_idx();
|
||||
|
||||
const int64_t idx_dim = hparams.indexer_head_size;
|
||||
const int64_t n_idx_h = hparams.indexer_n_head;
|
||||
const int64_t r = hparams.dsv4_compress_ratios[il];
|
||||
const int64_t n_kv = mctx_idx->get_n_kv();
|
||||
const int64_t idx_dim = hparams.indexer_head_size;
|
||||
const int64_t n_idx_h = hparams.indexer_n_head;
|
||||
const int64_t kpool = inp_kpool->kpool;
|
||||
const int64_t n_pool = inp_kpool->pool_cells->ne[0];
|
||||
const int64_t n_new = inp_kpool->n_new;
|
||||
|
||||
GGML_ASSERT(r > 0);
|
||||
GGML_ASSERT(hparams.dsv4_compress_ratios[il] == kpool);
|
||||
|
||||
const int64_t n_blocks = (n_kv + r - 1)/r;
|
||||
|
||||
// build_attn_qsa and the KQ mask need the tokens to divide evenly across the streams
|
||||
const int64_t n_stream = mctx_hyb->get_n_stream();
|
||||
GGML_ASSERT(n_tokens % n_stream == 0);
|
||||
const int64_t n_tps = n_tokens/n_stream;
|
||||
|
||||
// only the "which block is visible" half of the bias varies per block
|
||||
// the rest is the visible/not test the attention mask already carries, so upload the per-block half only: 1/ratio of the cells
|
||||
// alibi writes distances instead of a mask, so it opts out
|
||||
// the mask also holds an mrope rule for the query's own position, but only 2d image positions can differ there
|
||||
const bool blk_bias = kq_mask != nullptr &&
|
||||
kq_mask->ne[0] == n_kv && kq_mask->ne[1] == n_tps && kq_mask->ne[3] == n_stream &&
|
||||
!hparams.use_alibi;
|
||||
|
||||
// nothing above depends on the layer, so the layers sharing a ratio share one input set
|
||||
llm_graph_input_qsa * inp = nullptr;
|
||||
|
||||
const auto it = qsa_inps.find((uint32_t) r);
|
||||
if (it != qsa_inps.end()) {
|
||||
inp = it->second;
|
||||
} else {
|
||||
auto qsa = std::make_unique<llm_graph_input_qsa>(mctx_hyb, (uint32_t) r, blk_bias, cparams.causal_attn);
|
||||
|
||||
qsa->k_idxs = mctx_idx->build_input_k_idxs(ctx0, ubatch);
|
||||
qsa->cell_blk = ggml_new_tensor_2d(ctx0, GGML_TYPE_I32, n_kv, n_stream);
|
||||
qsa->blk_cells = ggml_new_tensor_2d(ctx0, GGML_TYPE_I32, r*n_blocks, n_stream);
|
||||
qsa->blk_pos = ggml_new_tensor_1d(ctx0, GGML_TYPE_I32, 4*n_blocks*n_stream);
|
||||
qsa->bias = ggml_new_tensor_3d(ctx0, GGML_TYPE_F32, blk_bias ? n_blocks : n_kv, n_tps, n_stream);
|
||||
|
||||
ggml_set_input(qsa->cell_blk);
|
||||
ggml_set_input(qsa->blk_cells);
|
||||
ggml_set_input(qsa->blk_pos);
|
||||
ggml_set_input(qsa->bias);
|
||||
|
||||
inp = qsa.get();
|
||||
res->add_input(std::move(qsa));
|
||||
qsa_inps.emplace((uint32_t) r, inp);
|
||||
}
|
||||
|
||||
// cached indexer keys are raw: pooling precedes norm and rotation, so apply neither
|
||||
// cache rows store raw key | pooled key: pooling precedes norm and rotation, so the raw key gets neither
|
||||
ggml_tensor * k_raw = build_lora_mm(model.layers[il].index_k_proj, cur);
|
||||
k_raw = ggml_reshape_3d(ctx0, k_raw, idx_dim, 1, n_tokens);
|
||||
cb(k_raw, "indexer_k_raw", il);
|
||||
|
||||
ggml_build_forward_expand(gf, mctx_idx->cpy_k(ctx0, k_raw, inp->k_idxs, il));
|
||||
ggml_tensor * pzero = ggml_fill(ctx0, ggml_new_tensor_2d(ctx0, GGML_TYPE_F32, idx_dim, n_tokens), 0.0f);
|
||||
ggml_tensor * packed = ggml_reshape_3d(ctx0, ggml_concat(ctx0, k_raw, pzero, 0), 2*idx_dim, 1, n_tokens);
|
||||
ggml_build_forward_expand(gf, mctx_idx->cpy_k(ctx0, packed, inp_kpool->k_idxs, il));
|
||||
|
||||
// one key head, so rows are contiguous. get_k gives [idx_dim, n_head_kv, n_kv, n_stream].
|
||||
ggml_tensor * k_all = mctx_idx->get_k(ctx0, il);
|
||||
k_all = ggml_view_3d(ctx0, k_all, idx_dim, n_kv, n_stream, k_all->nb[2], k_all->nb[3], 0);
|
||||
// the raw keys and the persistent pooled slots, see llama_memory_hybrid_idx::mem_idx_stale
|
||||
auto kpool_cache = mctx_hyb->get_kpool_access(ctx0, il, idx_dim);
|
||||
|
||||
// gathers per stream: blk_cells row s indexes stream s's own cells
|
||||
ggml_tensor * members = ggml_get_rows(ctx0, k_all, inp->blk_cells);
|
||||
members = ggml_reshape_4d(ctx0, members, idx_dim, r, n_blocks, n_stream);
|
||||
// pool only the blocks this ubatch completes or regroups
|
||||
ggml_tensor * rows = kpool_cache.gather_key_gate(ggml_reshape_1d(ctx0, inp_kpool->new_pool_idxs, kpool*n_new));
|
||||
rows = ggml_reshape_3d(ctx0, rows, idx_dim, kpool, n_new);
|
||||
|
||||
// mean over the block members; r is small, so summing slices beats a transpose plus sum_rows
|
||||
ggml_tensor * pooled = nullptr;
|
||||
for (int64_t i = 0; i < r; ++i) {
|
||||
ggml_tensor * slice = ggml_cont(ctx0,
|
||||
ggml_view_3d(ctx0, members, idx_dim, n_blocks, n_stream,
|
||||
members->nb[2], members->nb[3], i*members->nb[1]));
|
||||
pooled = pooled ? ggml_add(ctx0, pooled, slice) : slice;
|
||||
// mean over the members; kpool is small, so summing slices beats a transpose plus sum_rows
|
||||
ggml_tensor * pooled_new = nullptr;
|
||||
for (int64_t i = 0; i < kpool; ++i) {
|
||||
ggml_tensor * slice = ggml_view_2d(ctx0, rows, idx_dim, n_new, rows->nb[2], i*rows->nb[1]);
|
||||
pooled_new = pooled_new ? ggml_add(ctx0, pooled_new, slice) : ggml_cont(ctx0, slice);
|
||||
}
|
||||
pooled = ggml_scale(ctx0, pooled, 1.0f/(float) r);
|
||||
cb(pooled, "indexer_k_pooled", il);
|
||||
pooled_new = ggml_scale(ctx0, pooled_new, 1.0f/(float) kpool);
|
||||
pooled_new = build_norm(pooled_new, model.layers[il].index_k_norm, nullptr, LLM_NORM_RMS, il);
|
||||
|
||||
// count blocks along ne1: rms_norm launches gridDim.y = ne2, capped at 65535, and 262144/4 = 65536
|
||||
pooled = ggml_reshape_3d(ctx0, pooled, idx_dim, n_blocks*n_stream, 1);
|
||||
pooled = build_norm(pooled, model.layers[il].index_k_norm, nullptr, LLM_NORM_RMS, il);
|
||||
|
||||
// rope wants [n_dims, n_head, n_tokens]: lay every stream's blocks flat, split after.
|
||||
pooled = ggml_reshape_3d(ctx0, pooled, idx_dim, 1, n_blocks*n_stream);
|
||||
pooled = ggml_rope_multi(ctx0, pooled, inp->blk_pos, nullptr,
|
||||
pooled_new = ggml_reshape_3d(ctx0, pooled_new, idx_dim, 1, n_new);
|
||||
pooled_new = ggml_rope_multi(ctx0, pooled_new, inp_kpool->new_pool_pos, nullptr,
|
||||
n_rot, sections, rope_type, n_ctx_orig, freq_base, freq_scale,
|
||||
ext_factor, attn_factor, beta_fast, beta_slow);
|
||||
pooled = ggml_reshape_3d(ctx0, pooled, idx_dim, n_blocks, n_stream);
|
||||
pooled_new = ggml_reshape_2d(ctx0, pooled_new, idx_dim, n_new);
|
||||
cb(pooled_new, "indexer_pool_k_new", il);
|
||||
|
||||
ggml_tensor * pooled = nullptr;
|
||||
if (inp_kpool->cache_safe) {
|
||||
// write before the pool gather
|
||||
ggml_build_forward_expand(gf, kpool_cache.scatter_pooled(pooled_new, inp_kpool->new_pool_rep));
|
||||
pooled = kpool_cache.gather_pooled(inp_kpool->pool_cells);
|
||||
} else {
|
||||
// shared cells re-pool every pool, in layout order
|
||||
GGML_ASSERT(n_new < n_pool);
|
||||
ggml_tensor * pad = ggml_fill(ctx0, ggml_new_tensor_2d(ctx0, GGML_TYPE_F32, idx_dim, n_pool - n_new), 0.0f);
|
||||
pooled = ggml_concat(ctx0, pooled_new, pad, 1);
|
||||
}
|
||||
pooled = ggml_reshape_3d(ctx0, pooled, idx_dim, 1, n_pool);
|
||||
cb(pooled, "indexer_k", il);
|
||||
|
||||
ggml_tensor * q = build_lora_mm(model.layers[il].index_q_proj, cur);
|
||||
@@ -649,67 +685,73 @@ ggml_tensor * llama_model_qwen4exp::graph::build_qsa_top_k(
|
||||
ext_factor, attn_factor, beta_fast, beta_slow);
|
||||
cb(q, "indexer_q", il);
|
||||
|
||||
// rectify each head dot product before the sum, as in the DeepSeek lightning indexer
|
||||
// mul_mat matches ne[2], so the queries of stream s only meet the blocks of stream s
|
||||
ggml_tensor * score = ggml_mul_mat(ctx0, pooled,
|
||||
ggml_reshape_3d(ctx0, q, idx_dim, n_idx_h*n_tps, n_stream));
|
||||
score = ggml_reshape_4d(ctx0, score, n_blocks, n_idx_h, n_tps, n_stream);
|
||||
score = ggml_relu(ctx0, score);
|
||||
// the reference sums the rectified head scores unweighted, scaled by 1/sqrt(head_dim)
|
||||
// one product for all heads, then the heads are summed as slices, so nothing is transposed
|
||||
ggml_tensor * kq = ggml_mul_mat(ctx0,
|
||||
ggml_reshape_2d(ctx0, pooled, idx_dim, n_pool),
|
||||
ggml_reshape_2d(ctx0, q, idx_dim, n_idx_h*n_tokens)); // [n_pool, n_idx_h*n_tokens]
|
||||
kq = ggml_relu(ctx0, ggml_reshape_3d(ctx0, kq, n_pool, n_idx_h, n_tokens));
|
||||
|
||||
// the heads sit side by side on ne[1] and there are only a few of them
|
||||
ggml_tensor * summed = nullptr;
|
||||
ggml_tensor * score = nullptr;
|
||||
for (int64_t h = 0; h < n_idx_h; ++h) {
|
||||
ggml_tensor * slice = ggml_view_3d(ctx0, score, n_blocks, n_tps, n_stream,
|
||||
score->nb[2], score->nb[3], h*score->nb[1]);
|
||||
summed = summed ? ggml_add(ctx0, summed, slice) : ggml_cont(ctx0, slice);
|
||||
ggml_tensor * slice = ggml_view_2d(ctx0, kq, n_pool, n_tokens, kq->nb[2], h*kq->nb[1]);
|
||||
score = score ? ggml_add(ctx0, score, slice) : ggml_cont(ctx0, slice);
|
||||
}
|
||||
|
||||
score = summed;
|
||||
score = ggml_scale(ctx0, score, 1.0f/sqrtf((float) idx_dim));
|
||||
score = ggml_add(ctx0, score, inp_kpool->pool_mask); // [n_pool, n_tokens]
|
||||
cb(score, "indexer_score", il);
|
||||
|
||||
// one value per block, so it is cheaper to bias here than after the cells are expanded
|
||||
if (blk_bias) {
|
||||
score = ggml_add(ctx0, score, inp->bias);
|
||||
}
|
||||
|
||||
// every token of a block gets the block score; the budget is whole blocks, so top-k cuts on a block boundary
|
||||
ggml_tensor * expanded = ggml_get_rows(ctx0,
|
||||
ggml_cont(ctx0, ggml_permute(ctx0, score, 1, 0, 2, 3)), inp->cell_blk);
|
||||
expanded = ggml_cont(ctx0, ggml_permute(ctx0, expanded, 1, 0, 2, 3));
|
||||
|
||||
if (blk_bias) {
|
||||
// flash attention keeps the mask in f16; the scores are f32
|
||||
ggml_tensor * mask = kq_mask->type == GGML_TYPE_F32 ? kq_mask : ggml_cast(ctx0, kq_mask, GGML_TYPE_F32);
|
||||
expanded = ggml_add(ctx0, expanded, ggml_reshape_3d(ctx0, mask, n_kv, n_tps, n_stream));
|
||||
} else {
|
||||
expanded = ggml_add(ctx0, expanded, inp->bias);
|
||||
}
|
||||
cb(expanded, "indexer_score_tokens", il);
|
||||
|
||||
// the reference returns indexer_top_k + compress_ratio - 1: whole blocks plus the tail
|
||||
const int64_t width = std::min<int64_t>(n_kv, (int64_t) hparams.indexer_top_k + r - 1);
|
||||
|
||||
ggml_tensor * top_k = ggml_cont(ctx0, ggml_top_k(ctx0, expanded, width));
|
||||
|
||||
// build_attn_qsa reads [n_top_k, n_batch, 1, n_stream], matching the KQ mask.
|
||||
top_k = ggml_reshape_4d(ctx0, top_k, width, n_tps, 1, n_stream);
|
||||
const int64_t n_top_pool = std::min<int64_t>(n_pool, hparams.indexer_top_k / kpool);
|
||||
ggml_tensor * top_k = ggml_top_k(ctx0, score, n_top_pool); // [n_top_pool, n_tokens], unordered
|
||||
cb(top_k, "indexer_top_k", il);
|
||||
|
||||
return top_k;
|
||||
// the top blocks, then the incomplete tail with n_kv for missing cells
|
||||
ggml_tensor * sel_idx = ggml_get_rows(ctx0, inp_kpool->pool_idxs,
|
||||
ggml_reshape_1d(ctx0, top_k, n_top_pool*n_tokens)); // [kpool, n_top_pool*n_tokens]
|
||||
sel_idx = ggml_reshape_2d(ctx0, sel_idx, kpool*n_top_pool, n_tokens);
|
||||
sel_idx = ggml_concat(ctx0, sel_idx, inp_kpool->tail_idxs, 0);
|
||||
const int64_t n_sel = sel_idx->ne[0];
|
||||
GGML_ASSERT(n_sel == inp_kpool->n_sel);
|
||||
|
||||
// scatter zeros for the selected cells into an all -inf row, the extra row n_kv takes the sentinels
|
||||
// seeding from sel_idx ties the scatter storage lifetime to this layer
|
||||
const int64_t n_kv = inp_kpool->n_kv;
|
||||
|
||||
ggml_tensor * seed = ggml_cast(ctx0, ggml_view_1d(ctx0, sel_idx, 1, 0), GGML_TYPE_F32);
|
||||
|
||||
ggml_tensor * mask_seed = kq_mask->type == GGML_TYPE_F32 ? seed : ggml_cast(ctx0, seed, kq_mask->type);
|
||||
mask_seed = ggml_fill(ctx0, mask_seed, -INFINITY);
|
||||
ggml_tensor * mask_all = ggml_repeat_4d(ctx0, mask_seed, 1, n_kv + 1, n_tokens, 1);
|
||||
mask_all = ggml_reshape_3d(ctx0, mask_all, 1, n_kv + 1, n_tokens);
|
||||
|
||||
ggml_tensor * zero_seed = ggml_fill(ctx0, seed, 0.0f);
|
||||
ggml_tensor * zeros = ggml_repeat_4d(ctx0, zero_seed, 1, n_sel, n_tokens, 1);
|
||||
zeros = ggml_reshape_3d(ctx0, zeros, 1, n_sel, n_tokens);
|
||||
|
||||
ggml_tensor * sel = ggml_set_rows(ctx0, mask_all, zeros, ggml_reshape_3d(ctx0, sel_idx, n_sel, n_tokens, 1));
|
||||
|
||||
GGML_ASSERT(kq_mask->ne[0] == n_kv && kq_mask->ne[1]*kq_mask->ne[2]*kq_mask->ne[3] == n_tokens);
|
||||
const size_t row = sel->nb[2];
|
||||
sel = ggml_view_4d(ctx0, sel, n_kv, kq_mask->ne[1], kq_mask->ne[2], kq_mask->ne[3],
|
||||
row, row*kq_mask->ne[1], row*kq_mask->ne[1]*kq_mask->ne[2], 0);
|
||||
sel = ggml_add(ctx0, sel, kq_mask);
|
||||
cb(sel, "indexer_sel", il);
|
||||
|
||||
return sel;
|
||||
}
|
||||
|
||||
// Dense GQA self-attention restricted to the cells that top_k names.
|
||||
// The mask build below copies the MLA sparse path in llm_graph_context::build_attn.
|
||||
// Dense GQA self-attention over the cells that the QSA mask keeps.
|
||||
ggml_tensor * llama_model_qwen4exp::graph::build_attn_qsa(
|
||||
llm_graph_input_attn_kv * inp,
|
||||
ggml_tensor * q_cur,
|
||||
ggml_tensor * k_cur,
|
||||
ggml_tensor * v_cur,
|
||||
ggml_tensor * top_k,
|
||||
ggml_tensor * sel,
|
||||
int64_t n_sel,
|
||||
float kq_scale,
|
||||
int il) {
|
||||
// rotate q/k/v before they reach a quantized cache, as the dense path does. the indexer
|
||||
// has already scored with its own query in build_qsa_top_k, so top_k is unaffected.
|
||||
// has already scored with its own query in build_qsa_sel, so the selection is unaffected.
|
||||
if (inp->self_k_rot) {
|
||||
q_cur = llama_mul_mat_hadamard(ctx0, q_cur, inp->self_k_rot);
|
||||
k_cur = llama_mul_mat_hadamard(ctx0, k_cur, inp->self_k_rot);
|
||||
@@ -737,39 +779,16 @@ ggml_tensor * llama_model_qwen4exp::graph::build_attn_qsa(
|
||||
ggml_build_forward_expand(gf, mctx_cur->cpy_v(ctx0, v_cur, v_idxs, il));
|
||||
}
|
||||
|
||||
// the selection mask already carries the causal mask
|
||||
ggml_tensor * kq_mask = inp->get_kq_mask();
|
||||
|
||||
// prepare new kq mask - starts filled with -INFINITY
|
||||
ggml_tensor * kq_mask_all = ggml_fill(ctx0, kq_mask, -INFINITY);
|
||||
|
||||
// reshape KQ mask into tensor with rows of size 1:
|
||||
// [n_kv, n_batch, 1, n_stream] -> [1, n_kv, n_batch, n_stream]
|
||||
kq_mask_all = ggml_view_4d(ctx0, kq_mask_all, 1, kq_mask_all->ne[0], kq_mask_all->ne[1], kq_mask_all->ne[3], kq_mask_all->nb[0], kq_mask_all->nb[1], kq_mask_all->nb[2], 0);
|
||||
|
||||
// reshape top_k indices: [n_top_k, n_batch, 1, n_stream] -> [n_top_k, n_batch, n_stream, 1]
|
||||
ggml_tensor * top_k_3d = ggml_view_4d(ctx0, top_k, top_k->ne[0], top_k->ne[1], top_k->ne[3], 1, top_k->nb[1], top_k->nb[2], top_k->ne[3]*top_k->nb[3], 0);
|
||||
|
||||
// prepare zero-filled tensor with rows of size 1: [1, n_top_k, n_batch, n_stream]
|
||||
// this will be our source of zero values for unmasking top k mask elements
|
||||
ggml_tensor * zeros = ggml_new_tensor_4d(ctx0, GGML_TYPE_F32, 1, top_k_3d->ne[0], top_k_3d->ne[1], top_k_3d->ne[2]);
|
||||
zeros = ggml_fill(ctx0, zeros, 0.0f);
|
||||
|
||||
// modify KQ mask by unmasking elements that are in top_k indices
|
||||
// ggml_set_rows([1, n_kv, n_batch, n_stream], [1, n_top_k, n_batch, n_stream], [n_top_k, n_batch, n_stream, 1])
|
||||
ggml_tensor * kq_mask_top_k = ggml_set_rows(ctx0, kq_mask_all, zeros, top_k_3d);
|
||||
|
||||
// reshape to restore the original shape of KQ mask:
|
||||
// [1, n_kv, n_batch, n_stream] -> [n_kv, n_batch, 1, n_stream]
|
||||
kq_mask_top_k = ggml_view_4d(ctx0, kq_mask_top_k, kq_mask_top_k->ne[1], kq_mask_top_k->ne[2], 1, kq_mask_top_k->ne[3], kq_mask_top_k->nb[2], kq_mask_top_k->nb[3], kq_mask_top_k->nb[3], 0);
|
||||
|
||||
// combine with the original kq mask
|
||||
kq_mask_top_k = ggml_add(ctx0, kq_mask_top_k, kq_mask);
|
||||
ggml_tensor * mask = ggml_reshape_4d(ctx0, sel, kq_mask->ne[0], kq_mask->ne[1], kq_mask->ne[2], kq_mask->ne[3]);
|
||||
cb(mask, "kq_mask_qsa", il);
|
||||
|
||||
ggml_tensor * q = q_cur;
|
||||
ggml_tensor * k = mctx_cur->get_k(ctx0, il);
|
||||
ggml_tensor * v = mctx_cur->get_v(ctx0, il);
|
||||
|
||||
ggml_tensor * cur = build_attn_mha(q, k, v, nullptr, kq_mask_top_k, nullptr, nullptr, top_k->ne[0], kq_scale, il);
|
||||
ggml_tensor * cur = build_attn_mha(q, k, v, nullptr, mask, nullptr, nullptr, n_sel, kq_scale, il);
|
||||
cb(cur, "kqv_out", il);
|
||||
|
||||
// the rotation is its own inverse, so undo it on the value side of the output
|
||||
@@ -783,6 +802,7 @@ ggml_tensor * llama_model_qwen4exp::graph::build_attn_qsa(
|
||||
ggml_tensor * llama_model_qwen4exp::graph::build_layer_attn(
|
||||
llm_graph_input_attn_kv * inp,
|
||||
const llama_memory_hybrid_idx_context * mctx_hyb,
|
||||
llm_graph_input_kpool * inp_kpool,
|
||||
ggml_tensor * cur,
|
||||
ggml_tensor * inp_pos,
|
||||
int * sections,
|
||||
@@ -791,9 +811,9 @@ ggml_tensor * llama_model_qwen4exp::graph::build_layer_attn(
|
||||
GGML_ASSERT(n_embd_head == hparams.n_embd_head_k());
|
||||
|
||||
// indexer reads the same block input as q/k/v; no cache or no ratio means dense
|
||||
const bool qsa = mctx_hyb->get_idx() != nullptr && hparams.dsv4_compress_ratios[il] > 0;
|
||||
const bool qsa = inp_kpool != nullptr && hparams.dsv4_compress_ratios[il] > 0;
|
||||
|
||||
ggml_tensor * top_k = qsa ? build_qsa_top_k(mctx_hyb, cur, inp_pos, inp->get_kq_mask(), sections, il) : nullptr;
|
||||
ggml_tensor * sel = qsa ? build_qsa_sel(mctx_hyb, inp_kpool, cur, inp_pos, inp->get_kq_mask(), sections, il) : nullptr;
|
||||
|
||||
// Qwen3Next uses a single Q projection that outputs query + gate
|
||||
ggml_tensor * Qcur_full = build_lora_mm(model.layers[il].wq, cur, model.layers[il].wq_s); // [ (n_embd_head * 2) * n_head, n_tokens ]
|
||||
@@ -845,8 +865,8 @@ ggml_tensor * llama_model_qwen4exp::graph::build_layer_attn(
|
||||
|
||||
const float kq_scale = hparams.f_attention_scale == 0.0f ? 1.0f / sqrtf(float(n_embd_head)) : hparams.f_attention_scale;
|
||||
|
||||
if (top_k) {
|
||||
cur = build_attn_qsa(inp, Qcur, Kcur, Vcur, top_k, kq_scale, il);
|
||||
if (sel) {
|
||||
cur = build_attn_qsa(inp, Qcur, Kcur, Vcur, sel, inp_kpool->n_sel, kq_scale, il);
|
||||
} else {
|
||||
cur = build_attn(inp,
|
||||
nullptr, nullptr, nullptr,
|
||||
|
||||
Reference in New Issue
Block a user