ggml-vulkan: implement host buffer pinning for faster H2D uploads

Assisted-by: opencode
This commit is contained in:
2026-07-12 15:51:29 +02:00
parent 5e7f6271c0
commit 8a545c7820
+53 -1
View File
@@ -17917,11 +17917,63 @@ static ggml_backend_dev_t ggml_backend_vk_reg_get_device(ggml_backend_reg_t reg,
return devices[device];
}
static bool ggml_backend_vk_register_host_buffer(void * buffer, size_t size) {
if (getenv("GGML_CUDA_REGISTER_HOST") == nullptr && getenv("GGML_VK_REGISTER_HOST") == nullptr) {
return false;
}
bool success = false;
for (size_t i = 0; i < GGML_VK_MAX_DEVICES; i++) {
vk_device& device = vk_instance.devices[i];
if (!device || !device->external_memory_host) continue;
vk_buffer buf = ggml_vk_buffer_from_host_ptr(device, buffer, size);
if (!buf) {
continue;
}
std::lock_guard<std::shared_mutex> guard(device->pinned_memory_mutex);
device->pinned_memory.push_back(std::make_tuple(buffer, size, buf));
success = true;
}
return success;
}
static void ggml_backend_vk_unregister_host_buffer(void * buffer) {
for (size_t i = 0; i < GGML_VK_MAX_DEVICES; i++) {
vk_device& device = vk_instance.devices[i];
if (!device) continue;
std::lock_guard<std::shared_mutex> guard(device->pinned_memory_mutex);
for (auto it = device->pinned_memory.begin(); it != device->pinned_memory.end(); ) {
if (std::get<0>(*it) == buffer) {
vk_buffer buf = std::get<2>(*it);
ggml_vk_destroy_buffer(buf);
it = device->pinned_memory.erase(it);
break; // A buffer is registered once per device
} else {
++it;
}
}
}
}
static void * ggml_backend_vk_reg_get_proc_address(ggml_backend_reg_t reg, const char * name) {
UNUSED(reg);
if (strcmp(name, "ggml_backend_register_host_buffer") == 0) {
return (void *)ggml_backend_vk_register_host_buffer;
}
if (strcmp(name, "ggml_backend_unregister_host_buffer") == 0) {
return (void *)ggml_backend_vk_unregister_host_buffer;
}
return nullptr;
}
static const struct ggml_backend_reg_i ggml_backend_vk_reg_i = {
/* .get_name = */ ggml_backend_vk_reg_get_name,
/* .get_device_count = */ ggml_backend_vk_reg_get_device_count,
/* .get_device = */ ggml_backend_vk_reg_get_device,
/* .get_proc_address = */ NULL,
/* .get_proc_address = */ ggml_backend_vk_reg_get_proc_address,
};
ggml_backend_reg_t ggml_backend_vk_reg() {