mirror of
https://github.com/ggml-org/llama.cpp.git
synced 2026-09-28 17:07:31 -05:00
context : do not re-reserve the scheduler when toggling causal_attn (#28751)
* context : do not re-reserve the scheduler when toggling causal_attn `llama_context::set_causal_attn()` marks the scheduler to do a full re-reserve on every change of the flag. For vision inputs, this flag is flipped twice around each non-causal image chunk for Gemma models, resulting in two expensive `sched_reserve()` passes per image. This is especially slow for multi-image or video inputs. The cost of a re-reserve scales with context and ubatch configurations, so larger settings pay more per image (see table below). The re-reserve is unnecessary in this case because `causal_attn` only changes the values written to KQ mask, not tensor shapes or any other buffer sizes. Note: `causal_attn` is a graph reuse key (`llm_graph_params` via `cparams`), so a new graph is built regardless of `sched_need_reserve`, so this doesn't change the graph rebuilding behaviour. llama-server with gemma-4-26B-A4B Q4_0 + BF16 mmproj, 130-token images, cache_prompt=false, prompt_ms median of 3 (before -> after): | images | config | H200 before -> after | RTX 4090 before -> after | |-|-|-|-| | 1 | `-c 8192 -ub 512` | 134 -> 105 ms (1.27×) | 201 -> 119 ms (1.69×) | | 24 | `-c 8192 -ub 512` | 2278 -> 1562 ms (1.46×) | 3559 -> 1748 ms (2.04×) | | 24 | `-c 32768 -ub 2048` | 5379 -> 1584 ms (3.40×) | 13377 -> 1759 ms (7.61×) | Generated output remains identical before and after. * qwen4exp : make the indexer bias shape independent of causal_attn The block/cell bias path was selected on cparams.causal_attn, so the causal and non-causal graphs differed in tensor shapes and ops. With the re-reserve removed (previous commit), a runtime flip resulted in reallocating the compute buffers, which would fail under GGML_SCHED_NO_REALLOC. This commit selects the block path from the mask shape only, independent of causal_attn. causal_attn is instead passed to set_input_qsa. causal_attn is fixed per graph as it's part of the reuse key. Causal values are unchanged. Non-causal values now follow the reference rule, where every visible block competes on score and only unpooled cells are always selected. * context : state the causal_attn shape rule in the comment * cont : add TODOs --------- Co-authored-by: Georgi Gerganov <ggerganov@gmail.com>
This commit is contained in:
co-authored by
Georgi Gerganov
parent
0c6a6a7ce5
commit
ed7ac35e1e
@@ -1256,7 +1256,8 @@ void llama_context::set_causal_attn(bool value) {
|
||||
|
||||
cparams.causal_attn = value;
|
||||
|
||||
sched_need_reserve = true;
|
||||
// no scheduler reserve needed because graph shapes must not depend on causal_attn, a flip only rebuilds the graph
|
||||
//sched_need_reserve = true;
|
||||
}
|
||||
|
||||
bool llama_context::get_causal_attn() const {
|
||||
|
||||
@@ -277,7 +277,8 @@ void llama_memory_hybrid_idx::set_input_qsa(
|
||||
ggml_tensor * bias,
|
||||
const llama_ubatch * ubatch,
|
||||
uint32_t ratio,
|
||||
bool blk_bias) const {
|
||||
bool blk_bias,
|
||||
bool causal_attn) const {
|
||||
GGML_ASSERT(ratio > 0);
|
||||
GGML_ASSERT(get_mem_idx() != nullptr);
|
||||
|
||||
@@ -545,7 +546,7 @@ void llama_memory_hybrid_idx::set_input_qsa(
|
||||
|
||||
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 future cells
|
||||
// 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) {
|
||||
@@ -555,7 +556,7 @@ void llama_memory_hybrid_idx::set_input_qsa(
|
||||
}
|
||||
|
||||
// finite, so it can never meet a -inf and produce a nan
|
||||
cur_blk_bias[b] = bid_idx[b] >= tail_start ? 1e9f : 0.0f;
|
||||
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
|
||||
@@ -576,7 +577,10 @@ void llama_memory_hybrid_idx::set_input_qsa(
|
||||
if (!cells.is_empty(j) && cells.seq_has(j, seq_id)) {
|
||||
const int64_t idx = ranked ? rank[j] : cells.pos_get(j);
|
||||
|
||||
if (idx <= q) {
|
||||
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);
|
||||
}
|
||||
@@ -676,8 +680,9 @@ void llama_memory_hybrid_idx_context::set_input_qsa(
|
||||
ggml_tensor * bias,
|
||||
const llama_ubatch * ubatch,
|
||||
uint32_t ratio,
|
||||
bool blk_bias) const {
|
||||
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);
|
||||
mem->set_input_qsa(cell_blk, blk_cells, blk_pos, bias, ubatch, ratio, blk_bias, causal_attn);
|
||||
}
|
||||
|
||||
@@ -12,6 +12,8 @@
|
||||
// llama_memory_hybrid plus a third cache with one indexer key per token, for block-sparse attention (qwen4exp QSA)
|
||||
// the indexer is a side buffer over the attention cells: same size, padding, streams and slots, so cell j is one token in both
|
||||
|
||||
// TODO: this memory module is pending complete reimplementation - do not use for model other than Qwen4
|
||||
|
||||
class llama_memory_hybrid_idx : public llama_memory_hybrid {
|
||||
public:
|
||||
llama_memory_hybrid_idx(
|
||||
@@ -83,9 +85,10 @@ public:
|
||||
// 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) const;
|
||||
bool blk_bias, bool causal_attn) const;
|
||||
|
||||
private:
|
||||
// forget seq_id (all of it if seq_id < 0) in every cache at once, so a failed restore cannot leave the caches out of step
|
||||
@@ -143,7 +146,7 @@ public:
|
||||
|
||||
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) const;
|
||||
bool blk_bias, bool causal_attn) const;
|
||||
|
||||
private:
|
||||
const llama_memory_hybrid_idx * mem = nullptr;
|
||||
|
||||
+12
-6
@@ -6,6 +6,9 @@
|
||||
#include <algorithm>
|
||||
#include <cinttypes>
|
||||
|
||||
// [TAG_QWEN4_REIMPLEMENT]
|
||||
// TODO: this graph implementation is pending complete reimplementation - do not use it as a reference
|
||||
|
||||
// bad metadata must be catchable: GGML_ASSERT aborts the whole process
|
||||
static void qwen4exp_require_nonzero(const llama_model_loader & ml, llm_kv kid, uint32_t value) {
|
||||
if (value == 0) {
|
||||
@@ -489,13 +492,13 @@ ggml_tensor * llama_model_qwen4exp::graph::build_norm_gated(
|
||||
// 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 {
|
||||
public:
|
||||
llm_graph_input_qsa(const llama_memory_hybrid_idx_context * mctx, uint32_t ratio, bool blk_bias) :
|
||||
mctx(mctx), ratio(ratio), blk_bias(blk_bias) {}
|
||||
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;
|
||||
|
||||
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);
|
||||
mctx->set_input_qsa(cell_blk, blk_cells, blk_pos, bias, ubatch, ratio, blk_bias, causal_attn);
|
||||
}
|
||||
|
||||
bool can_reuse(const llm_graph_params & params) override {
|
||||
@@ -537,6 +540,9 @@ public:
|
||||
|
||||
// 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;
|
||||
};
|
||||
|
||||
ggml_tensor * llama_model_qwen4exp::graph::build_qsa_top_k(
|
||||
@@ -564,11 +570,11 @@ ggml_tensor * llama_model_qwen4exp::graph::build_qsa_top_k(
|
||||
|
||||
// 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 and non-causal keeps future cells, so both opt out
|
||||
// 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 &&
|
||||
cparams.causal_attn && !hparams.use_alibi;
|
||||
!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;
|
||||
@@ -577,7 +583,7 @@ ggml_tensor * llama_model_qwen4exp::graph::build_qsa_top_k(
|
||||
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);
|
||||
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);
|
||||
|
||||
Reference in New Issue
Block a user