vulkan : reuse descriptor sets when bindings are constant (#29280)

* vulkan : reuse descriptor sets when bindings are constant

* vulkan : bump buffer_destroy_count before destroying the buffer
This commit is contained in:
Andrew Lee
2026-09-29 08:54:55 +02:00
committed by GitHub
parent 139997d8e7
commit 0bc845d356
3 changed files with 41 additions and 3 deletions
+30 -3
View File
@@ -263,9 +263,36 @@ inline void ggml_vk_dispatch_pipeline(ggml_backend_vk_context* ctx, vk_context&
GGML_ASSERT(pipeline->parameter_count == descriptor_buffer_infos.size());
GGML_ASSERT(pipeline->push_constant_size == push_constant_size(push_constants));
vk::DescriptorSet& descriptor_set = ctx->descriptor_sets[ctx->descriptor_set_idx++];
vk::WriteDescriptorSet write_descriptor_set{ descriptor_set, 0, 0, pipeline->parameter_count, vk::DescriptorType::eStorageBuffer, nullptr, descriptor_buffer_infos.begin() };
ctx->device->device.updateDescriptorSets({ write_descriptor_set }, {});
const uint32_t descriptor_set_idx = ctx->descriptor_set_idx++;
vk::DescriptorSet& descriptor_set = ctx->descriptor_sets[descriptor_set_idx];
// a new buffer can get the handle of a destroyed one, so drop all cached bindings after any destroy
const uint64_t destroy_count = ctx->device->buffer_destroy_count.load(std::memory_order_acquire);
if (ctx->descriptor_set_bindings_destroy_count != destroy_count) {
for (auto & b : ctx->descriptor_set_bindings) {
b.clear();
}
ctx->descriptor_set_bindings_destroy_count = destroy_count;
}
// skip the write if this set already holds these bindings from the last graph
std::vector<vk::DescriptorBufferInfo> & bindings = ctx->descriptor_set_bindings[descriptor_set_idx];
bool same = !ctx->device->disable_descriptor_reuse && bindings.size() == descriptor_buffer_infos.size();
if (same) {
size_t i = 0;
for (const vk::DescriptorBufferInfo & info : descriptor_buffer_infos) {
const vk::DescriptorBufferInfo & prev = bindings[i++];
if (prev.buffer != info.buffer || prev.offset != info.offset || prev.range != info.range) {
same = false;
break;
}
}
}
if (!same) {
vk::WriteDescriptorSet write_descriptor_set{ descriptor_set, 0, 0, pipeline->parameter_count, vk::DescriptorType::eStorageBuffer, nullptr, descriptor_buffer_infos.begin() };
ctx->device->device.updateDescriptorSets({ write_descriptor_set }, {});
bindings.assign(descriptor_buffer_infos.begin(), descriptor_buffer_infos.end());
}
subctx->s->buffer->buf.pushConstants(pipeline->layout, vk::ShaderStageFlagBits::eCompute, 0, push_constant_size(push_constants), push_constant_data(push_constants));
subctx->s->buffer->buf.bindPipeline(vk::PipelineBindPoint::eCompute, pipeline->pipeline);
+7
View File
@@ -1025,6 +1025,8 @@ struct vk_device_struct {
ggml_backend_buffer_type buffer_type;
bool disable_fusion;
bool disable_descriptor_reuse;
std::atomic<uint64_t> buffer_destroy_count {};
bool disable_host_visible_vidmem;
bool allow_sysmem_fallback;
bool disable_graph_optimize;
@@ -1067,6 +1069,8 @@ struct vk_buffer_struct {
}
VK_LOG_DEBUG("~vk_buffer_struct(" << buffer << ", " << size << ")");
// bump before destroying, so a thread that sees the buffer gone also sees the new count
device->buffer_destroy_count.fetch_add(1, std::memory_order_release);
device->device.freeMemory(device_memory);
device->device.destroyBuffer(buffer);
}
@@ -1278,6 +1282,9 @@ struct ggml_backend_vk_context {
std::vector<vk::DescriptorPool> descriptor_pools;
std::vector<vk::DescriptorSet> descriptor_sets;
// last bindings written to each set; descriptor_sets is append-only so an index always names the same set
std::vector<std::vector<vk::DescriptorBufferInfo>> descriptor_set_bindings;
uint64_t descriptor_set_bindings_destroy_count {};
uint32_t descriptor_set_idx {};
uint32_t pipeline_descriptor_set_requirements {};
+4
View File
@@ -881,6 +881,7 @@ void ggml_pipeline_allocate_descriptor_sets(ggml_backend_vk_context * ctx) {
pool_idx++;
}
ctx->descriptor_set_bindings.resize(ctx->descriptor_sets.size());
}
static vk_command_buffer* ggml_vk_create_cmd_buffer(vk_device& device, vk_command_pool& p) {
@@ -4885,6 +4886,8 @@ vk_device ggml_vk_get_device(size_t idx) {
device->disable_fusion = getenv("GGML_VK_DISABLE_FUSION") != nullptr;
device->disable_descriptor_reuse = getenv("GGML_VK_DISABLE_DESCRIPTOR_REUSE") != nullptr;
device->add_rms_fusion = !device->disable_fusion &&
device->subgroup_arithmetic &&
device->vendor_id != VK_VENDOR_ID_INTEL;
@@ -12846,6 +12849,7 @@ void ggml_vk_cleanup(ggml_backend_vk_context * ctx) {
}
ctx->descriptor_pools.clear();
ctx->descriptor_sets.clear();
ctx->descriptor_set_bindings.clear();
ctx->compute_cmd_pool.destroy(ctx->device->device);
if (ctx->device->async_use_transfer_queue) {