feat: add configurable Qwen cache types and early cache scheduling (#2045)

This commit is contained in:
leejet
2026-09-25 01:51:01 +08:00
committed by GitHub
parent 4dfe8f5d45
commit 740c7ae193
11 changed files with 185 additions and 61 deletions
+8 -1
View File
@@ -623,7 +623,11 @@ ggml_tensor* ggml_ext_attention_ext(ggml_context* ctx,
bool skip_reshape,
bool flash_attn,
float kv_scale,
bool sage_attn) { // avoid overflow
bool sage_attn,
bool* used_flash_attn) { // avoid overflow
if (used_flash_attn != nullptr) {
*used_flash_attn = false;
}
int64_t L_q;
int64_t L_k;
int64_t C;
@@ -755,6 +759,9 @@ ggml_tensor* ggml_ext_attention_ext(ggml_context* ctx,
if (can_use_flash_attn) {
kqv = build_kqv(q, k, v, mask);
if (kqv != nullptr) {
if (used_flash_attn != nullptr) {
*used_flash_attn = true;
}
kqv = ggml_view_4d(ctx,
kqv,
d_head,
+6 -5
View File
@@ -217,11 +217,12 @@ ggml_tensor* ggml_ext_attention_ext(ggml_context* ctx,
ggml_tensor* k,
ggml_tensor* v,
int64_t n_head,
ggml_tensor* mask = nullptr,
bool skip_reshape = false,
bool flash_attn = false,
float kv_scale = 1.0f,
bool sage_attn = false);
ggml_tensor* mask = nullptr,
bool skip_reshape = false,
bool flash_attn = false,
float kv_scale = 1.0f,
bool sage_attn = false,
bool* used_flash_attn = nullptr);
ggml_tensor* ggml_ext_layer_norm(ggml_context* ctx,
ggml_tensor* x,
+12 -6
View File
@@ -21,11 +21,12 @@ ggml_tensor* ggml_ext_attention_ext(GGMLRunnerContext* ctx,
ggml_tensor* mask,
bool skip_reshape,
bool flash_attn,
float kv_scale) {
float kv_scale,
bool* used_flash_attn) {
if (ctx->attn_scale > 0.f) {
kv_scale = ctx->attn_scale;
}
return ggml_ext_attention_ext(ctx->ggml_ctx, ctx->backend, q, k, v, n_head, mask, skip_reshape, flash_attn, kv_scale, ctx->sage_attn_enabled);
return ggml_ext_attention_ext(ctx->ggml_ctx, ctx->backend, q, k, v, n_head, mask, skip_reshape, flash_attn, kv_scale, ctx->sage_attn_enabled, used_flash_attn);
}
void GGMLRunner::alloc_params_ctx() {
@@ -515,9 +516,10 @@ GGMLRunner::~GGMLRunner() {
free_params_ctx();
}
GGMLRunnerContext GGMLRunner::get_context() {
GGMLRunnerContext GGMLRunner::get_context(ggml_cgraph* graph) {
GGMLRunnerContext runner_ctx;
runner_ctx.ggml_ctx = compute_ctx;
runner_ctx.graph = graph;
runner_ctx.backend = runtime_backend;
runner_ctx.flash_attn_enabled = flash_attn_enabled;
runner_ctx.sage_attn_enabled = sage_attn_enabled;
@@ -532,8 +534,8 @@ GGMLRunnerContext GGMLRunner::get_context() {
runner_ctx.get_cache_tensor = [this](const std::string& name) {
return this->get_cache_tensor_by_name(name);
};
runner_ctx.cache_tensor = [this](const std::string& name, ggml_tensor* tensor) {
this->cache(name, tensor);
runner_ctx.cache_tensor = [this, graph](const std::string& name, ggml_tensor* tensor) {
this->cache(name, tensor, graph);
};
runner_ctx.set_backend_tensor_data = [this](ggml_tensor* tensor, const void* data) {
this->set_backend_tensor_data(tensor, data);
@@ -575,7 +577,7 @@ ggml_tensor* GGMLRunner::to_backend(ggml_tensor* tensor) {
}
}
void GGMLRunner::cache(const std::string name, ggml_tensor* tensor) {
void GGMLRunner::cache(const std::string name, ggml_tensor* tensor, ggml_cgraph* graph) {
if (tensor != nullptr && tensor->view_src != nullptr) {
tensor = ggml_cont(compute_ctx, tensor);
}
@@ -583,6 +585,10 @@ void GGMLRunner::cache(const std::string name, ggml_tensor* tensor) {
ggml_set_output(tensor);
}
cache_.stage(name, tensor);
if (graph != nullptr && tensor != nullptr) {
// Schedule the cache output here so its source can be reused before graph end.
ggml_build_forward_expand(graph, tensor);
}
}
std::optional<sd::Tensor<float>> GGMLRunner::compute(get_graph_cb_t get_graph,
+15 -6
View File
@@ -67,6 +67,7 @@ struct WeightAdapter {
struct GGMLRunnerContext {
ggml_backend_t backend = nullptr;
ggml_context* ggml_ctx = nullptr;
ggml_cgraph* graph = nullptr;
bool flash_attn_enabled = false;
bool sage_attn_enabled = false;
float linear_scale = 0.f;
@@ -102,6 +103,12 @@ struct GGMLRunnerContext {
return get_cache_tensor(name);
}
void expand_graph(ggml_tensor* tensor) const {
if (graph != nullptr && tensor != nullptr) {
ggml_build_forward_expand(graph, tensor);
}
}
void persist_cache_tensor(const std::string& name, ggml_tensor* tensor) const {
if (!cache_tensor || tensor == nullptr) {
return;
@@ -122,10 +129,11 @@ ggml_tensor* ggml_ext_attention_ext(GGMLRunnerContext* ctx,
ggml_tensor* k,
ggml_tensor* v,
int64_t n_head,
ggml_tensor* mask = nullptr,
bool skip_reshape = false,
bool flash_attn = false,
float kv_scale = 1.f);
ggml_tensor* mask = nullptr,
bool skip_reshape = false,
bool flash_attn = false,
float kv_scale = 1.f,
bool* used_flash_attn = nullptr);
struct GGMLRunner {
private:
@@ -289,7 +297,8 @@ public:
virtual ~GGMLRunner();
virtual GGMLRunnerContext get_context();
// Binding a graph schedules cache outputs at registration instead of graph end.
virtual GGMLRunnerContext get_context(ggml_cgraph* graph = nullptr);
void reset_compute_ctx();
@@ -324,7 +333,7 @@ public:
ggml_tensor* to_backend(ggml_tensor* tensor);
void cache(const std::string name, ggml_tensor* tensor);
void cache(const std::string name, ggml_tensor* tensor, ggml_cgraph* graph = nullptr);
ggml_tensor* get_cache_tensor_by_name(const std::string& name) {
return cache_.get(name);