diff --git a/ggml/src/ggml-vulkan/ggml-vulkan-common.h b/ggml/src/ggml-vulkan/ggml-vulkan-common.h index 5f95312a76..125e7608a3 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan-common.h +++ b/ggml/src/ggml-vulkan/ggml-vulkan-common.h @@ -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 & 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); diff --git a/ggml/src/ggml-vulkan/ggml-vulkan-types.h b/ggml/src/ggml-vulkan/ggml-vulkan-types.h index bc5de0e875..ee50f51a38 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan-types.h +++ b/ggml/src/ggml-vulkan/ggml-vulkan-types.h @@ -1025,6 +1025,8 @@ struct vk_device_struct { ggml_backend_buffer_type buffer_type; bool disable_fusion; + bool disable_descriptor_reuse; + std::atomic 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 descriptor_pools; std::vector descriptor_sets; + // last bindings written to each set; descriptor_sets is append-only so an index always names the same set + std::vector> descriptor_set_bindings; + uint64_t descriptor_set_bindings_destroy_count {}; uint32_t descriptor_set_idx {}; uint32_t pipeline_descriptor_set_requirements {}; diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index 29077dfc8e..38ec1f9e82 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -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) {