feat: add Qwen Image 2.1 prefix KV cache (#2035)

This commit is contained in:
leejet
2026-09-23 23:08:21 +08:00
committed by GitHub
parent e6281b6318
commit 2dc7f5408a
9 changed files with 237 additions and 81 deletions
+2
View File
@@ -71,6 +71,8 @@ struct AnimaDiffusionExtra {
struct QwenImage21DiffusionExtra {
const sd::Tensor<int32_t>* image_slots = nullptr;
// Nonzero IDs identify immutable prefix inputs within one sampling run.
uint64_t prefix_id = 0;
};
struct WanDiffusionExtra {
+165 -60
View File
@@ -121,6 +121,18 @@ namespace Qwen {
}
};
struct QwenImage21PrefixCache {
enum class Mode {
NONE,
STORE,
REUSE
};
Mode mode = Mode::NONE;
std::string name;
std::string cut_group;
int64_t prefix_length = 0;
};
class QwenImage21ZeroCenterRMSNorm : public RMSNorm {
public:
using RMSNorm::RMSNorm;
@@ -160,7 +172,7 @@ namespace Qwen {
}
}
ggml_tensor* forward(GGMLRunnerContext* ctx, ggml_tensor* x, ggml_tensor* pe, const std::vector<QwenImage21Segment>& segments, const std::vector<ggml_tensor*>& masks) {
ggml_tensor* forward(GGMLRunnerContext* ctx, ggml_tensor* x, ggml_tensor* pe, const std::vector<QwenImage21Segment>& segments, const std::vector<ggml_tensor*>& masks, const QwenImage21PrefixCache& cache) {
int64_t heads = x->ne[0] / dim_head;
auto project = [&](const char* name) {
auto h = std::dynamic_pointer_cast<Linear>(blocks[name])->forward(ctx, x);
@@ -173,14 +185,36 @@ namespace Qwen {
k = std::dynamic_pointer_cast<RMSNorm>(blocks["norm_k"])->forward(ctx, k);
q = Rope::apply_rope(ctx->ggml_ctx, q, pe);
k = Rope::apply_rope(ctx->ggml_ctx, k, pe);
if (cache.mode == QwenImage21PrefixCache::Mode::STORE) {
auto persist = [&](ggml_tensor* tensor, int axis, const char* name) {
auto part = ggml_ext_slice(ctx->ggml_ctx, tensor, axis, 0, cache.prefix_length);
auto copy = ggml_new_tensor(ctx->ggml_ctx, GGML_TYPE_F32, 4, part->ne);
copy = ggml_cpy(ctx->ggml_ctx, part, copy);
// Keep the copy in this layer's segment so graph cuts do not
// retain or recompute the full-sequence K/V in the final segment.
sd::ggml_graph_cut::mark_graph_cut(copy, cache.cut_group, name);
ctx->persist_cache_tensor(cache.name + "." + name, copy);
};
persist(k, 1, "k");
persist(v, 2, "v");
}
ggml_tensor* result = nullptr;
for (size_t i = 0; i < segments.size(); ++i) {
const auto& segment = segments[i];
auto sq = ggml_ext_slice(ctx->ggml_ctx, q, 1, segment.start, segment.end);
auto sk = ggml_ext_slice(ctx->ggml_ctx, k, 1, 0, segment.end);
auto sv = ggml_ext_slice(ctx->ggml_ctx, v, 2, 0, segment.end);
auto out = ggml_ext_attention_ext(ctx, sq, sk, sv, heads, masks[i], true, ctx->flash_attn_enabled);
result = result == nullptr ? out : ggml_concat(ctx->ggml_ctx, result, out, 1);
if (cache.mode == QwenImage21PrefixCache::Mode::REUSE) {
auto prefix_k = ctx->load_cache_tensor(cache.name + ".k");
auto prefix_v = ctx->load_cache_tensor(cache.name + ".v");
GGML_ASSERT(prefix_k != nullptr && prefix_v != nullptr);
k = ggml_concat(ctx->ggml_ctx, prefix_k, k, 1);
v = ggml_concat(ctx->ggml_ctx, prefix_v, v, 2);
result = ggml_ext_attention_ext(ctx, q, k, v, heads, nullptr, true, ctx->flash_attn_enabled);
} else {
for (size_t i = 0; i < segments.size(); ++i) {
const auto& segment = segments[i];
auto sq = ggml_ext_slice(ctx->ggml_ctx, q, 1, segment.start, segment.end);
auto sk = ggml_ext_slice(ctx->ggml_ctx, k, 1, 0, segment.end);
auto sv = ggml_ext_slice(ctx->ggml_ctx, v, 2, 0, segment.end);
auto out = ggml_ext_attention_ext(ctx, sq, sk, sv, heads, masks[i], true, ctx->flash_attn_enabled);
result = result == nullptr ? out : ggml_concat(ctx->ggml_ctx, result, out, 1);
}
}
auto to_out = std::dynamic_pointer_cast<Linear>(blocks["to_out.0"]);
if (sd_backend_is(ctx->backend, "Vulkan") || sd_backend_is(ctx->backend, "ROCm")) {
@@ -219,13 +253,14 @@ namespace Qwen {
return ggml_concat(ctx, prefix, target, 1);
}
ggml_tensor* forward(GGMLRunnerContext* ctx, ggml_tensor* x, const std::vector<ggml_tensor*>& modulation, ggml_tensor* pe, const QwenImage21Layout& layout, const std::vector<ggml_tensor*>& masks) {
auto h = std::dynamic_pointer_cast<LayerNorm>(blocks["img_norm1"])->forward(ctx, x);
h = modulate(ctx->ggml_ctx, h, modulation[0], layout.prefix_length);
h = std::dynamic_pointer_cast<QwenImage21Attention>(blocks["attn"])->forward(ctx, h, pe, layout.segments, masks);
x = ggml_add(ctx->ggml_ctx, x, modulate(ctx->ggml_ctx, h, modulation[1], layout.prefix_length, true));
h = std::dynamic_pointer_cast<LayerNorm>(blocks["img_norm2"])->forward(ctx, x);
h = modulate(ctx->ggml_ctx, h, modulation[2], layout.prefix_length);
ggml_tensor* forward(GGMLRunnerContext* ctx, ggml_tensor* x, const std::vector<ggml_tensor*>& modulation, ggml_tensor* pe, const QwenImage21Layout& layout, const std::vector<ggml_tensor*>& masks, const QwenImage21PrefixCache& cache) {
const int64_t prefix_length = cache.mode == QwenImage21PrefixCache::Mode::REUSE ? 0 : layout.prefix_length;
auto h = std::dynamic_pointer_cast<LayerNorm>(blocks["img_norm1"])->forward(ctx, x);
h = modulate(ctx->ggml_ctx, h, modulation[0], prefix_length);
h = std::dynamic_pointer_cast<QwenImage21Attention>(blocks["attn"])->forward(ctx, h, pe, layout.segments, masks, cache);
x = ggml_add(ctx->ggml_ctx, x, modulate(ctx->ggml_ctx, h, modulation[1], prefix_length, true));
h = std::dynamic_pointer_cast<LayerNorm>(blocks["img_norm2"])->forward(ctx, x);
h = modulate(ctx->ggml_ctx, h, modulation[2], prefix_length);
ggml_tensor* gate;
auto fused = blocks.find("img_mlp.gate_up");
if (fused != blocks.end()) {
@@ -239,7 +274,7 @@ namespace Qwen {
}
h = ggml_mul(ctx->ggml_ctx, h, ggml_silu(ctx->ggml_ctx, gate));
h = std::dynamic_pointer_cast<Linear>(blocks["img_mlp.out"])->forward(ctx, h);
return ggml_add(ctx->ggml_ctx, x, modulate(ctx->ggml_ctx, h, modulation[3], layout.prefix_length, true));
return ggml_add(ctx->ggml_ctx, x, modulate(ctx->ggml_ctx, h, modulation[3], prefix_length, true));
}
};
@@ -261,7 +296,7 @@ namespace Qwen {
}
}
ggml_tensor* forward(GGMLRunnerContext* ctx, ggml_tensor* x, ggml_tensor* timestep, ggml_tensor* context, const std::vector<ggml_tensor*>& refs, ggml_tensor* pe, const QwenImage21Layout& layout, const std::vector<ggml_tensor*>& masks) {
ggml_tensor* forward(GGMLRunnerContext* ctx, ggml_tensor* x, ggml_tensor* timestep, ggml_tensor* context, const std::vector<ggml_tensor*>& refs, ggml_tensor* pe, const QwenImage21Layout& layout, const std::vector<ggml_tensor*>& masks, const QwenImage21PrefixCache& cache) {
auto time = ggml_concat(ctx->ggml_ctx, timestep, ggml_ext_zeros_like(ctx->ggml_ctx, timestep), 0);
// Runtime flow timesteps already use the [0, 1000] scale.
time = ggml_ext_timestep_embedding(ctx->ggml_ctx, time, 256, 10000, 1.f);
@@ -269,27 +304,37 @@ namespace Qwen {
time = ggml_silu(ctx->ggml_ctx, time);
auto modulation = std::dynamic_pointer_cast<Linear>(blocks["modulation.1"])->forward(ctx, time);
auto mod = ggml_ext_chunk(ctx->ggml_ctx, modulation, 4, 0);
auto text = std::dynamic_pointer_cast<QwenImage21TextProjection>(blocks["txt_in"])->forward(ctx, context);
auto img_in = std::dynamic_pointer_cast<Linear>(blocks["img_in"]);
ggml_tensor* joint = nullptr;
for (const auto& segment : layout.segments) {
ggml_tensor* h;
if (segment.image_index < 0) {
h = ggml_ext_slice(ctx->ggml_ctx, text, 1, segment.context_start,
segment.context_start + segment.end - segment.start);
} else {
auto image = segment.image_index == static_cast<int>(refs.size()) ? x : refs[segment.image_index];
h = img_in->forward(ctx, DiT::patchify(ctx->ggml_ctx, image, 1, 1));
if (cache.mode == QwenImage21PrefixCache::Mode::REUSE) {
joint = img_in->forward(ctx, DiT::patchify(ctx->ggml_ctx, x, 1, 1));
} else {
auto text = std::dynamic_pointer_cast<QwenImage21TextProjection>(blocks["txt_in"])->forward(ctx, context);
for (const auto& segment : layout.segments) {
ggml_tensor* h;
if (segment.image_index < 0) {
h = ggml_ext_slice(ctx->ggml_ctx, text, 1, segment.context_start,
segment.context_start + segment.end - segment.start);
} else {
auto image = segment.image_index == static_cast<int>(refs.size()) ? x : refs[segment.image_index];
h = img_in->forward(ctx, DiT::patchify(ctx->ggml_ctx, image, 1, 1));
}
joint = joint == nullptr ? h : ggml_concat(ctx->ggml_ctx, joint, h, 1);
}
joint = joint == nullptr ? h : ggml_concat(ctx->ggml_ctx, joint, h, 1);
}
sd::ggml_graph_cut::mark_graph_cut(joint, "qwen_image_2_1.prelude", "joint");
for (int i = 0; i < config.num_layers; ++i) {
auto block = std::dynamic_pointer_cast<QwenImage21TransformerBlock>(blocks["transformer_blocks." + std::to_string(i)]);
joint = block->forward(ctx, joint, mod, pe, layout, masks);
sd::ggml_graph_cut::mark_graph_cut(joint, "qwen_image_2_1.transformer_blocks." + std::to_string(i), "joint");
const std::string layer = "transformer_blocks." + std::to_string(i);
auto layer_cache = cache;
layer_cache.name = cache.name + "." + std::to_string(i);
layer_cache.cut_group = "qwen_image_2_1." + layer;
auto block = std::dynamic_pointer_cast<QwenImage21TransformerBlock>(blocks[layer]);
joint = block->forward(ctx, joint, mod, pe, layout, masks, layer_cache);
sd::ggml_graph_cut::mark_graph_cut(joint, layer_cache.cut_group, "joint");
}
if (cache.mode != QwenImage21PrefixCache::Mode::REUSE) {
joint = ggml_ext_slice(ctx->ggml_ctx, joint, 1, layout.prefix_length, joint->ne[1]);
}
joint = ggml_ext_slice(ctx->ggml_ctx, joint, 1, layout.prefix_length, joint->ne[1]);
auto scale = std::dynamic_pointer_cast<Linear>(blocks["norm_out.linear"])->forward(ctx, ggml_ext_chunk(ctx->ggml_ctx, time, 2, 1)[0]);
joint = std::dynamic_pointer_cast<LayerNorm>(blocks["norm_out.norm"])->forward(ctx, joint);
joint = ggml_mul(ctx->ggml_ctx, joint, ggml_scale_bias(ctx->ggml_ctx, scale, 1.f, 1.f));
@@ -303,11 +348,18 @@ namespace Qwen {
QwenImage21Model model;
std::vector<float> pe_data;
std::vector<sd::Tensor<float>> mask_data;
bool prefix_cache_enabled = true;
bool prefix_cache_disabled = false;
QwenImage21Runner(ggml_backend_t backend, const String2TensorStorage& weights, const std::string& prefix, std::shared_ptr<RunnerWeightManager> weight_manager = nullptr)
QwenImage21Runner(ggml_backend_t backend, const String2TensorStorage& weights, const std::string& prefix, std::shared_ptr<RunnerWeightManager> weight_manager = nullptr, const char* model_args = nullptr)
: DiffusionModelRunner(backend, prefix, weight_manager),
config(QwenImage21Config::detect_from_weights(weights, prefix)),
model(config) {
for (const auto& [key, value] : parse_key_value_args(model_args, "model arg")) {
if (key == "qwen_image_2_1_prefix_cache" && !parse_strict_bool(value, prefix_cache_enabled)) {
LOG_WARN("ignoring invalid Qwen Image 2.1 model arg '%s=%s'", key.c_str(), value.c_str());
}
}
model.init(params_ctx, weights, prefix);
}
@@ -317,6 +369,22 @@ namespace Qwen {
model.get_param_tensors(tensors, prefix);
}
bool has_prefix_cache(const QwenImage21PrefixCache& cache) {
for (int i = 0; i < config.num_layers; ++i) {
const auto name = cache.name + "." + std::to_string(i);
auto k = get_cache_tensor_by_name(name + ".k");
auto v = get_cache_tensor_by_name(name + ".v");
if (k == nullptr || v == nullptr || k->type != GGML_TYPE_F32 || v->type != GGML_TYPE_F32 ||
k->ne[0] != config.head_dim || k->ne[1] != cache.prefix_length ||
k->ne[2] != config.hidden_size / config.head_dim || k->ne[3] != 1 ||
v->ne[0] != config.head_dim || v->ne[1] != config.hidden_size / config.head_dim ||
v->ne[2] != cache.prefix_length || v->ne[3] != 1) {
return false;
}
}
return true;
}
sd::Tensor<float> compute(int n_threads, const DiffusionParams& inputs) override {
const auto& x = tensor_or_empty(inputs.x);
const auto& context = tensor_or_empty(inputs.context);
@@ -345,38 +413,75 @@ namespace Qwen {
LOG_ERROR("%s", error.what());
return {};
}
pe_data = Rope::embed_nd(layout.positions, 1, 10000.f, config.axes_dim);
mask_data.clear();
for (const auto& segment : layout.segments) {
sd::Tensor<float> mask;
if (segment.image_index < 0) {
mask = sd::Tensor<float>::zeros({segment.end, segment.end - segment.start});
for (int64_t q = segment.start; q < segment.end; ++q) {
for (int64_t k = q + 1; k < segment.end; ++k) {
mask[k + segment.end * (q - segment.start)] = -INFINITY;
if (!runner_started()) {
prefix_cache_disabled = false;
}
QwenImage21PrefixCache cache;
if (prefix_cache_enabled && !prefix_cache_disabled && extra != nullptr && extra->prefix_id != 0 && layout.prefix_length > 0) {
cache.name = "qwen_image_2_1.prefix." + std::to_string(extra->prefix_id);
cache.prefix_length = layout.prefix_length;
cache.mode = has_prefix_cache(cache) ? QwenImage21PrefixCache::Mode::REUSE : QwenImage21PrefixCache::Mode::STORE;
}
auto run = [&](const QwenImage21PrefixCache& active_cache) {
const bool cached = active_cache.mode == QwenImage21PrefixCache::Mode::REUSE;
const auto first_position = layout.positions.begin() + (cached ? layout.prefix_length : 0);
pe_data = Rope::embed_nd(std::vector<std::vector<float>>(first_position, layout.positions.end()), 1, 10000.f, config.axes_dim);
mask_data.clear();
if (!cached) {
for (const auto& segment : layout.segments) {
sd::Tensor<float> mask;
if (segment.image_index < 0) {
mask = sd::Tensor<float>::zeros({segment.end, segment.end - segment.start});
for (int64_t q = segment.start; q < segment.end; ++q) {
for (int64_t k = q + 1; k < segment.end; ++k) {
mask[k + segment.end * (q - segment.start)] = -INFINITY;
}
}
}
mask_data.push_back(std::move(mask));
}
}
mask_data.push_back(std::move(mask));
}
auto build = [&]() {
auto graph = new_graph_custom(QWEN_IMAGE_GRAPH_SIZE * 2);
auto pe = ggml_new_tensor_4d(compute_ctx, GGML_TYPE_F32, 2, 2, config.head_dim / 2, layout.positions.size());
set_backend_tensor_data(pe, pe_data.data());
std::vector<ggml_tensor*> masks, ref_inputs;
for (const auto& mask : mask_data) {
masks.push_back(mask.empty() ? nullptr : make_input(mask));
}
for (const auto& ref : refs) {
ref_inputs.push_back(make_input(ref));
}
auto ctx = get_context();
auto out = model.forward(&ctx, make_input(x), make_input(*inputs.timesteps), make_input(context),
ref_inputs, pe, layout, masks);
ggml_build_forward_expand(graph, out);
return graph;
auto build = [&]() {
auto graph = new_graph_custom(QWEN_IMAGE_GRAPH_SIZE * 2);
auto pe = ggml_new_tensor_4d(compute_ctx, GGML_TYPE_F32, 2, 2, config.head_dim / 2,
layout.positions.size() - (cached ? layout.prefix_length : 0));
set_backend_tensor_data(pe, pe_data.data());
std::vector<ggml_tensor*> masks, ref_inputs;
for (const auto& mask : mask_data) {
masks.push_back(mask.empty() ? nullptr : make_input(mask));
}
if (!cached) {
for (const auto& ref : refs) {
ref_inputs.push_back(make_input(ref));
}
}
auto ctx = get_context();
auto out = model.forward(&ctx, make_input(x), make_input(*inputs.timesteps), cached ? nullptr : make_input(context),
ref_inputs, pe, layout, masks, active_cache);
ggml_build_forward_expand(graph, out);
return graph;
};
return restore_trailing_singleton_dims(GGMLRunner::compute(build, n_threads, false), x.dim());
};
return restore_trailing_singleton_dims(GGMLRunner::compute(build, n_threads, false), x.dim());
auto result = run(cache);
if (result.empty() && last_compute_status() == GGML_STATUS_ALLOC_FAILED &&
(cache.mode != QwenImage21PrefixCache::Mode::NONE || !cache_.empty())) {
// The failed graph has ended before persistent inputs are released.
free_cache_ctx_and_buffer();
prefix_cache_disabled = true;
LOG_WARN("Qwen Image 2.1: insufficient memory for prefix caching; retrying without it for this sampling run");
return run(QwenImage21PrefixCache{});
}
if (!result.empty() && cache.mode == QwenImage21PrefixCache::Mode::STORE) {
if (!has_prefix_cache(cache)) {
free_cache_ctx_and_buffer();
prefix_cache_disabled = true;
LOG_WARN("Qwen Image 2.1: incomplete prefix cache; disabling it for this sampling run");
} else {
LOG_DEBUG("Qwen Image 2.1: cached prefix %" PRIu64 " (%" PRId64 " tokens)", extra->prefix_id, layout.prefix_length);
}
}
return result;
}
};
}