diff --git a/ggml/src/ggml-backend-impl.h b/ggml/src/ggml-backend-impl.h index 9c56ec30c5..529768d3d3 100644 --- a/ggml/src/ggml-backend-impl.h +++ b/ggml/src/ggml-backend-impl.h @@ -83,6 +83,9 @@ extern "C" { GGML_API ggml_backend_buffer_t ggml_backend_multi_buffer_alloc_buffer(ggml_backend_buffer_t * buffers, size_t n_buffers); GGML_API bool ggml_backend_buffer_is_multi_buffer(ggml_backend_buffer_t buffer); GGML_API void ggml_backend_multi_buffer_set_usage(ggml_backend_buffer_t buffer, enum ggml_backend_buffer_usage usage); + // resolve the physical sub-buffer whose memory range contains addr (NULL if none); sub-buffers have a + // usable context/get_base, the multi-buffer wrapper does not, so backends need this to find the real buffer + GGML_API ggml_backend_buffer_t ggml_backend_multi_buffer_get_buffer(ggml_backend_buffer_t buffer, const void * addr); // // Backend (meta) diff --git a/ggml/src/ggml-backend-meta.cpp b/ggml/src/ggml-backend-meta.cpp index fe58ea3bb7..57f07b59df 100644 --- a/ggml/src/ggml-backend-meta.cpp +++ b/ggml/src/ggml-backend-meta.cpp @@ -1255,10 +1255,24 @@ static enum ggml_status ggml_backend_meta_buffer_init_tensor_impl(ggml_backend_m } if (t_ij->view_src != nullptr) { t_ij->data = (char *) t_ij->view_src->data + t_ij->view_offs; - } else if (simple_buf != nullptr) { + if (t_ij->view_src->buffer != nullptr) { + t_ij->buffer = t_ij->view_src->buffer; + } + } else if (simple_buf != nullptr && !ggml_backend_buffer_is_multi_buffer(simple_buf)) { + // single contiguous buffer: mirror the offset. A multi-buffer has no usable base, so leave the + // tensor for the gallocr to place (alloc_ctx_tensors assigns each tensor a real sub-buffer). t_ij->data = (char *) ggml_backend_buffer_get_base(simple_buf) + size_t(tensor->data) - size_t(ggml_backend_buffer_get_base(tensor->buffer)); } + // Backends require the physical buffer, not the multi-buffer wrapper (its context is unusable). + // Resolve it from the tensor's data address -- works for any tensor, not just views. + if (t_ij->buffer != nullptr && t_ij->data != nullptr + && ggml_backend_buffer_is_multi_buffer(t_ij->buffer)) { + ggml_backend_buffer_t sub = ggml_backend_multi_buffer_get_buffer(t_ij->buffer, t_ij->data); + if (sub != nullptr) { + t_ij->buffer = sub; + } + } t_ij->extra = tensor->extra; for (int i = 0; i < GGML_MAX_SRC; i++) { t_ij->src[i] = tensor->src[i]; diff --git a/ggml/src/ggml-backend.cpp b/ggml/src/ggml-backend.cpp index 3d6310f3ff..82b1103229 100644 --- a/ggml/src/ggml-backend.cpp +++ b/ggml/src/ggml-backend.cpp @@ -725,6 +725,20 @@ bool ggml_backend_buffer_is_multi_buffer(ggml_backend_buffer_t buffer) { return buffer->iface.free_buffer == ggml_backend_multi_buffer_free_buffer; } +ggml_backend_buffer_t ggml_backend_multi_buffer_get_buffer(ggml_backend_buffer_t buffer, const void * addr) { + GGML_ASSERT(ggml_backend_buffer_is_multi_buffer(buffer)); + ggml_backend_multi_buffer_context * ctx = (ggml_backend_multi_buffer_context *) buffer->context; + for (size_t i = 0; i < ctx->n_buffers; i++) { + ggml_backend_buffer_t sub = ctx->buffers[i]; + const char * base = (const char *) ggml_backend_buffer_get_base(sub); + const size_t size = ggml_backend_buffer_get_size(sub); + if ((const char *) addr >= base && (const char *) addr < base + size) { + return sub; + } + } + return NULL; +} + void ggml_backend_multi_buffer_set_usage(ggml_backend_buffer_t buffer, enum ggml_backend_buffer_usage usage) { GGML_ASSERT(buffer); GGML_ASSERT(ggml_backend_buffer_is_multi_buffer(buffer)); diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index c1d86aaac5..00ec7dfa0c 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -67,10 +67,12 @@ typedef struct VkPhysicalDeviceCooperativeMatrixDecodeVectorFeaturesNV { #include #include #include +#include #if defined(_MSC_VER) # define NOMINMAX 1 # include +# include // _aligned_malloc / _aligned_free # define YIELD() YieldProcessor() #elif defined(__clang__) || defined(__GNUC__) # if defined(__x86_64__) ||defined(__i386__) @@ -787,6 +789,7 @@ struct vk_device_struct { uint64_t suballocation_block_size; uint64_t min_imported_host_pointer_alignment; bool external_memory_host {}; + bool external_semaphore {}; bool fp16; bool bf16; bool pipeline_robustness; @@ -2338,6 +2341,10 @@ struct ggml_backend_vk_context { vk_device device; + bool comm_active = false; + vk::Semaphore comm_prog_sem = VK_NULL_HANDLE; + uint64_t comm_prog_val = 0; + size_t semaphore_idx, event_idx; ggml_vk_garbage_collector gc; size_t prealloc_size_x, prealloc_size_y, prealloc_size_split_k, prealloc_size_add_rms_partials, prealloc_size_add_rms_partials_offset; @@ -6271,6 +6278,8 @@ static vk_device ggml_vk_get_device(size_t idx) { device->memory_priority = true; } else if (strcmp("VK_EXT_external_memory_host", properties.extensionName) == 0) { device->external_memory_host = true; + } else if (strcmp("VK_KHR_external_semaphore_fd", properties.extensionName) == 0) { + device->external_semaphore = true; #if defined(VK_EXT_shader_64bit_indexing) } else if (strcmp("VK_EXT_shader_64bit_indexing", properties.extensionName) == 0) { device->shader_64b_indexing = true; @@ -6622,6 +6631,10 @@ static vk_device ggml_vk_get_device(size_t idx) { if (device->external_memory_host) { device_extensions.push_back("VK_EXT_external_memory_host"); } + if (device->external_semaphore) { + device_extensions.push_back("VK_KHR_external_semaphore"); + device_extensions.push_back("VK_KHR_external_semaphore_fd"); + } #if defined(VK_EXT_shader_64bit_indexing) VkPhysicalDeviceShader64BitIndexingFeaturesEXT shader_64bit_indexing_features {}; @@ -15809,6 +15822,11 @@ static bool ggml_vk_build_graph(ggml_backend_vk_context * ctx, ggml_cgraph * cgr if (submit || last_node) { ggml_vk_ctx_end(compute_ctx); + if (last_node && ctx->comm_active && !compute_ctx->seqs.empty()) { + ctx->comm_prog_val++; + compute_ctx->seqs.back().back().signal_semaphores.push_back({ ctx->comm_prog_sem, ctx->comm_prog_val }); + } + // TODO probably it'd be better to pass a exit_node flag to ggml_vk_compute_forward if (last_node) { compute_ctx->exit_tensor_idx = node_idx_begin; @@ -16400,9 +16418,17 @@ static bool ggml_backend_vk_cpy_tensor_async(ggml_backend_t backend_src, ggml_ba if (ggml_backend_buffer_is_vk(src->buffer)) { ggml_backend_vk_buffer_context * src_buf_ctx = (ggml_backend_vk_buffer_context *)src->buffer->context; - // Async copy only works within the same device + // Different devices have no peer-to-peer path, so stage through host memory. + // Flush the source backend so its data is ready, then run the blocking staged + // copy here and return true. This skips the scheduler's extra full synchronize + // of the destination backend (ggml_vk_buffer_copy already blocks until the write + // to the destination lands). if (src_buf_ctx->dev_buffer->device != dst_buf->device) { - return false; + ggml_backend_synchronize(backend_src); + ggml_vk_buffer_copy(dst_buf, vk_tensor_offset(dst) + dst->view_offs, + src_buf_ctx->dev_buffer, vk_tensor_offset(src) + src->view_offs, + ggml_nbytes(src)); + return true; } vk_context compute_ctx = ggml_vk_get_compute_ctx(ctx); @@ -17139,6 +17165,11 @@ static ggml_status ggml_backend_vk_graph_compute(ggml_backend_t backend, ggml_cg ggml_vk_sync_buffers(ctx, compute_ctx); } + if (ctx->comm_active) { + compute_ctx = ggml_vk_get_compute_ctx(ctx); + compute_ctx->s->wait_semaphores.push_back({ ctx->comm_prog_sem, ctx->comm_prog_val }); + } + // Submit after enough work has accumulated, to overlap CPU cmdbuffer generation with GPU execution. // Estimate the amount of compute work using flops, and submit every 200 GFLOP // (and scaled down based on total graph flops, so smaller models submit earlier). @@ -18844,11 +18875,883 @@ static ggml_backend_dev_t ggml_backend_vk_reg_get_device(ggml_backend_reg_t reg, return devices[device]; } +struct ggml_backend_vk_comm_context { + std::vector backends; + std::vector vkctx; + std::vector device; + size_t align = 4096; + size_t cap = 0; + bool fast = false; + std::vector prog; + std::vector> peer_prog; + std::vector last_reduce; + std::vector cmd_pool; + std::vector pool_max_val; + PFN_vkGetSemaphoreFdKHR pGetSemFd = nullptr; + PFN_vkImportSemaphoreFdKHR pImportSemFd = nullptr; + std::vector host_ptr; + std::vector> host_buf; + std::vector tmp_buffer; + std::vector tmp_tensor; + ggml_context * tctx = nullptr; + bool ring_ok = false; + std::vector up; + std::vector> peer_up; + std::vector up_val; + std::vector cmd_pool_xfer; + std::vector xfer_pool_max_val; + uint64_t ring_round = 0; + std::vector ring_view; + std::vector up16_buffer; + std::vector dn16_buffer; + std::vector up16_tensor; + std::vector dn16_tensor; + bool proxy = false; + std::vector pxy; + std::vector pxy_val; + struct bridge_task { vk::Device wdev; vk::Semaphore wsem; uint64_t wval; vk::Device sdev; vk::Semaphore ssem; uint64_t sval; }; + std::vector> bridge_q; + std::mutex bridge_mtx; + std::thread proxy_thread; + std::atomic proxy_stop{false}; +}; + +static void ggml_vk_comm_proxy_loop(ggml_backend_vk_comm_context * comm) { + while (!comm->proxy_stop.load(std::memory_order_relaxed)) { + bool progress = false; + for (size_t q = 0; q < comm->bridge_q.size(); q++) { + for (;;) { + ggml_backend_vk_comm_context::bridge_task t; + { + std::lock_guard lk(comm->bridge_mtx); + if (comm->bridge_q[q].empty()) { break; } + t = comm->bridge_q[q].front(); + if (t.wdev.getSemaphoreCounterValue(t.wsem) < t.wval) { break; } + comm->bridge_q[q].pop_front(); + } + vk::SemaphoreSignalInfo si{ t.ssem, t.sval }; + t.sdev.signalSemaphore(si); + progress = true; + } + } + if (!progress) { std::this_thread::yield(); } + } +} + +static vk::Semaphore ggml_vk_create_export_timeline(vk_device & device) { + vk::ExportSemaphoreCreateInfo esci{ vk::ExternalSemaphoreHandleTypeFlagBits::eOpaqueFd }; + vk::SemaphoreTypeCreateInfo stci{ vk::SemaphoreType::eTimeline, 0 }; + stci.pNext = &esci; + vk::SemaphoreCreateInfo sci{}; + sci.pNext = &stci; + return device->device.createSemaphore(sci); +} + +static vk::Semaphore ggml_vk_create_plain_timeline(vk_device & device) { + vk::SemaphoreTypeCreateInfo stci{ vk::SemaphoreType::eTimeline, 0 }; + vk::SemaphoreCreateInfo sci{}; + sci.pNext = &stci; + return device->device.createSemaphore(sci); +} + +static vk::Semaphore ggml_vk_import_timeline(ggml_backend_vk_comm_context * comm, vk_device & dst_dev, + vk_device & src_dev, vk::Semaphore src_sem) { + int fd = -1; + VkSemaphoreGetFdInfoKHR gi{}; + gi.sType = VK_STRUCTURE_TYPE_SEMAPHORE_GET_FD_INFO_KHR; + gi.semaphore = src_sem; + gi.handleType = VK_EXTERNAL_SEMAPHORE_HANDLE_TYPE_OPAQUE_FD_BIT; + comm->pGetSemFd((VkDevice) src_dev->device, &gi, &fd); + + vk::SemaphoreTypeCreateInfo stci{ vk::SemaphoreType::eTimeline, 0 }; + vk::SemaphoreCreateInfo sci{}; + sci.pNext = &stci; + vk::Semaphore dst = dst_dev->device.createSemaphore(sci); + + VkImportSemaphoreFdInfoKHR isi{}; + isi.sType = VK_STRUCTURE_TYPE_IMPORT_SEMAPHORE_FD_INFO_KHR; + isi.semaphore = dst; + isi.handleType = VK_EXTERNAL_SEMAPHORE_HANDLE_TYPE_OPAQUE_FD_BIT; + isi.fd = fd; + comm->pImportSemFd((VkDevice) dst_dev->device, &isi); + return dst; +} + +static bool ggml_vk_comm_opaque_fd_supported(ggml_backend_vk_comm_context * comm) { + const size_t n = comm->device.size(); + uint8_t driver_uuid[VK_UUID_SIZE]; + for (size_t i = 0; i < n; i++) { + vk::PhysicalDeviceIDProperties id; + vk::PhysicalDeviceProperties2 p2; + p2.pNext = &id; + comm->device[i]->physical_device.getProperties2(&p2); + if (i == 0) { + memcpy(driver_uuid, id.driverUUID.data(), VK_UUID_SIZE); + } else if (memcmp(driver_uuid, id.driverUUID.data(), VK_UUID_SIZE) != 0) { + return false; + } + + vk::SemaphoreTypeCreateInfo stci{ vk::SemaphoreType::eTimeline, 0 }; + vk::PhysicalDeviceExternalSemaphoreInfo esi{ vk::ExternalSemaphoreHandleTypeFlagBits::eOpaqueFd }; + esi.pNext = &stci; + vk::ExternalSemaphoreProperties esp = comm->device[i]->physical_device.getExternalSemaphoreProperties(esi); + const vk::ExternalSemaphoreFeatureFlags need = + vk::ExternalSemaphoreFeatureFlagBits::eExportable | vk::ExternalSemaphoreFeatureFlagBits::eImportable; + if ((esp.externalSemaphoreFeatures & need) != need) { + return false; + } + } + return true; +} + +static void * ggml_backend_vk_comm_init(ggml_backend_t * backends, size_t n_backends) { + ggml_backend_vk_comm_context * comm = new ggml_backend_vk_comm_context; + comm->backends.assign(backends, backends + n_backends); + comm->vkctx.resize(n_backends); + comm->device.resize(n_backends); + bool ok = (n_backends >= 2); + for (size_t i = 0; i < n_backends; i++) { + comm->vkctx[i] = (ggml_backend_vk_context *) backends[i]->context; + comm->device[i] = comm->vkctx[i]->device; + ok = ok && comm->device[i]->external_memory_host && comm->device[i]->external_semaphore; + comm->align = std::max(comm->align, (size_t) comm->device[i]->min_imported_host_pointer_alignment); + } + comm->fast = ok; + comm->host_ptr.resize(n_backends, nullptr); + comm->host_buf.resize(n_backends); + comm->tmp_buffer.resize(n_backends, nullptr); + comm->tmp_tensor.resize(n_backends, nullptr); + for (size_t k = 0; k < n_backends; k++) { + comm->host_buf[k].resize(n_backends); + } + comm->up16_buffer.resize(n_backends, nullptr); + comm->dn16_buffer.resize(n_backends, nullptr); + comm->up16_tensor.resize(n_backends, nullptr); + comm->dn16_tensor.resize(n_backends, nullptr); + comm->ring_view.resize(n_backends, nullptr); + const ggml_init_params ip = { ggml_tensor_overhead() * (4 * n_backends + 8), nullptr, true }; + comm->tctx = ggml_init(ip); + + if (comm->fast) { + comm->pGetSemFd = (PFN_vkGetSemaphoreFdKHR) comm->device[0]->device.getProcAddr("vkGetSemaphoreFdKHR"); + comm->pImportSemFd = (PFN_vkImportSemaphoreFdKHR) comm->device[0]->device.getProcAddr("vkImportSemaphoreFdKHR"); + comm->fast = comm->pGetSemFd && comm->pImportSemFd; + } + if (comm->fast) { + comm->prog.resize(n_backends); + comm->peer_prog.assign(n_backends, std::vector(n_backends)); + comm->last_reduce.assign(n_backends, 0); + comm->cmd_pool.resize(n_backends); + comm->pool_max_val.assign(n_backends, 0); + comm->up.resize(n_backends); + comm->peer_up.assign(n_backends, std::vector(n_backends)); + comm->up_val.assign(n_backends, 0); + comm->cmd_pool_xfer.resize(n_backends); + comm->xfer_pool_max_val.assign(n_backends, 0); + comm->ring_ok = true; + for (size_t i = 0; i < n_backends; i++) { + comm->cmd_pool[i].init(comm->device[i], &comm->device[i]->compute_queue); + comm->cmd_pool_xfer[i].init(comm->device[i], &comm->device[i]->transfer_queue); + if (comm->device[i]->single_queue || + comm->device[i]->transfer_queue.queue == comm->device[i]->compute_queue.queue) { + comm->ring_ok = false; + } + } + // Decide proxy vs native cross-device sync before creating the progress timelines: the proxy path only + // signals/reads them locally, so it must not request exportable handles. Some devices (e.g. llvmpipe, or + // RADV on older Mesa) cannot create exportable timeline semaphores at all, which would otherwise abort + // init here instead of falling back to the proxy. + comm->proxy = getenv("GGML_VK_COMM_PROXY") != nullptr; + if (!comm->proxy && !ggml_vk_comm_opaque_fd_supported(comm)) { + comm->proxy = true; + GGML_LOG_INFO("ggml_vulkan: cross-device OPAQUE_FD timeline import unsupported; using portable CPU-proxy sync\n"); + } + for (size_t i = 0; i < n_backends; i++) { + comm->prog[i] = comm->proxy ? ggml_vk_create_plain_timeline(comm->device[i]) + : ggml_vk_create_export_timeline(comm->device[i]); + comm->up[i] = comm->proxy ? ggml_vk_create_plain_timeline(comm->device[i]) + : ggml_vk_create_export_timeline(comm->device[i]); + } + if (!comm->proxy) { + try { + for (size_t j = 0; j < n_backends; j++) { + for (size_t i = 0; i < n_backends; i++) { + if (i == j) { continue; } + comm->peer_prog[i][j] = ggml_vk_import_timeline(comm, comm->device[i], comm->device[j], comm->prog[j]); + comm->peer_up[i][j] = ggml_vk_import_timeline(comm, comm->device[i], comm->device[j], comm->up[j]); + } + } + } catch (const vk::SystemError &) { + comm->proxy = true; + GGML_LOG_INFO("ggml_vulkan: comm OPAQUE_FD timeline import failed; using portable CPU-proxy sync\n"); + } + } + if (comm->proxy) { + comm->pxy.resize(n_backends); + comm->pxy_val.assign(n_backends, 0); + comm->bridge_q.resize(n_backends); + for (size_t i = 0; i < n_backends; i++) { + vk::SemaphoreTypeCreateInfo stci{ vk::SemaphoreType::eTimeline, 0 }; + vk::SemaphoreCreateInfo sci{}; + sci.pNext = &stci; + comm->pxy[i] = comm->device[i]->device.createSemaphore(sci); + } + comm->proxy_thread = std::thread(ggml_vk_comm_proxy_loop, comm); + } + for (size_t i = 0; i < n_backends; i++) { + comm->vkctx[i]->comm_prog_sem = comm->prog[i]; + comm->vkctx[i]->comm_prog_val = 0; + comm->vkctx[i]->comm_active = true; + } + } + return comm; +} + +static void * ggml_vk_comm_aligned_alloc(size_t alignment, size_t size) { +#if defined(_MSC_VER) || defined(__MINGW32__) + return _aligned_malloc(size, alignment); +#else + return std::aligned_alloc(alignment, size); +#endif +} + +static void ggml_vk_comm_aligned_free(void * ptr) { +#if defined(_MSC_VER) || defined(__MINGW32__) + _aligned_free(ptr); +#else + std::free(ptr); +#endif +} + +static void ggml_backend_vk_comm_free(void * comm_ctx) { + ggml_backend_vk_comm_context * comm = static_cast(comm_ctx); + for (size_t i = 0; i < comm->vkctx.size(); i++) { + ggml_backend_synchronize(comm->backends[i]); + comm->vkctx[i]->comm_active = false; + comm->vkctx[i]->comm_prog_sem = VK_NULL_HANDLE; + } + if (comm->proxy_thread.joinable()) { + comm->proxy_stop.store(true, std::memory_order_relaxed); + comm->proxy_thread.join(); + } + for (size_t i = 0; i < comm->pxy.size(); i++) { + if (comm->pxy[i]) { + comm->device[i]->device.destroySemaphore(comm->pxy[i]); + } + } + for (size_t i = 0; i < comm->prog.size(); i++) { + if (comm->prog[i]) { + comm->device[i]->device.destroySemaphore(comm->prog[i]); + } + if (i < comm->up.size() && comm->up[i]) { + comm->device[i]->device.destroySemaphore(comm->up[i]); + } + for (size_t j = 0; j < comm->peer_prog[i].size(); j++) { + if (comm->peer_prog[i][j]) { + comm->device[i]->device.destroySemaphore(comm->peer_prog[i][j]); + } + if (i < comm->peer_up.size() && j < comm->peer_up[i].size() && comm->peer_up[i][j]) { + comm->device[i]->device.destroySemaphore(comm->peer_up[i][j]); + } + } + } + for (size_t i = 0; i < comm->cmd_pool.size(); i++) { + comm->cmd_pool[i].destroy(comm->device[i]->device); + } + for (size_t i = 0; i < comm->cmd_pool_xfer.size(); i++) { + comm->cmd_pool_xfer[i].destroy(comm->device[i]->device); + } + for (auto & row : comm->host_buf) { + for (auto & b : row) { + b.reset(); + } + } + for (void * p : comm->host_ptr) { + ggml_vk_comm_aligned_free(p); + } + for (ggml_backend_buffer_t b : comm->tmp_buffer) { + if (b) { + ggml_backend_buffer_free(b); + } + } + for (ggml_backend_buffer_t b : comm->up16_buffer) { + if (b) { + ggml_backend_buffer_free(b); + } + } + for (ggml_backend_buffer_t b : comm->dn16_buffer) { + if (b) { + ggml_backend_buffer_free(b); + } + } + if (comm->tctx) { + ggml_free(comm->tctx); + } + delete comm; +} + +static bool ggml_backend_vk_comm_ensure(ggml_backend_vk_comm_context * comm, size_t nbytes) { + // Reuse the current buffers when they fit and are not grossly oversized. We deliberately SHRINK when the + // request is much smaller than the current cap (e.g. the prefill->decode transition): the imported external + // host buffers are made visible across devices on every timeline-semaphore signal, so an oversized `cap` + // left over from a large prefill stalls every small (decode) all-reduce in proportion to its size. + constexpr size_t shrink_slack = 4; + if (nbytes <= comm->cap && comm->cap <= nbytes * shrink_slack) { + return true; + } + const size_t n = comm->backends.size(); + const size_t newcap = (nbytes + comm->align - 1) & ~(comm->align - 1); + + // Buffers may still be referenced by the previous (async) all-reduce; wait for it before freeing them. + for (size_t i = 0; i < n; i++) { + if (comm->last_reduce[i] == 0) { + continue; + } + vk::SemaphoreWaitInfo wi; + wi.semaphoreCount = 1; + wi.pSemaphores = &comm->prog[i]; + wi.pValues = &comm->last_reduce[i]; + (void) comm->device[i]->device.waitSemaphores(wi, UINT64_MAX); + } + + for (size_t k = 0; k < n; k++) { + for (size_t i = 0; i < n; i++) { + comm->host_buf[k][i].reset(); + } + ggml_vk_comm_aligned_free(comm->host_ptr[k]); + comm->host_ptr[k] = nullptr; + } + for (size_t i = 0; i < n; i++) { + if (comm->tmp_buffer[i]) { + ggml_backend_buffer_free(comm->tmp_buffer[i]); + comm->tmp_buffer[i] = nullptr; + } + if (comm->up16_buffer[i]) { + ggml_backend_buffer_free(comm->up16_buffer[i]); + comm->up16_buffer[i] = nullptr; + } + if (comm->dn16_buffer[i]) { + ggml_backend_buffer_free(comm->dn16_buffer[i]); + comm->dn16_buffer[i] = nullptr; + } + } + + const size_t slotcap = 4 * newcap; + for (size_t k = 0; k < n; k++) { + comm->host_ptr[k] = ggml_vk_comm_aligned_alloc(comm->align, slotcap); + if (!comm->host_ptr[k]) { + return false; + } + for (size_t i = 0; i < n; i++) { + comm->host_buf[k][i] = ggml_vk_buffer_from_host_ptr(comm->device[i], comm->host_ptr[k], slotcap); + if (!comm->host_buf[k][i]) { + return false; + } + } + } + for (size_t i = 0; i < n; i++) { + ggml_backend_buffer_type_t bt = ggml_backend_get_default_buffer_type(comm->backends[i]); + comm->tmp_buffer[i] = ggml_backend_buft_alloc_buffer(bt, newcap); + if (!comm->tmp_buffer[i]) { + return false; + } + if (!comm->tmp_tensor[i]) { + comm->tmp_tensor[i] = ggml_new_tensor_1d(comm->tctx, GGML_TYPE_F32, 1); + } + comm->tmp_tensor[i]->buffer = comm->tmp_buffer[i]; + comm->tmp_tensor[i]->data = ggml_backend_buffer_get_base(comm->tmp_buffer[i]); + comm->up16_buffer[i] = ggml_backend_buft_alloc_buffer(bt, newcap); + comm->dn16_buffer[i] = ggml_backend_buft_alloc_buffer(bt, newcap); + if (!comm->up16_buffer[i] || !comm->dn16_buffer[i]) { + return false; + } + if (!comm->up16_tensor[i]) { + comm->up16_tensor[i] = ggml_new_tensor_1d(comm->tctx, GGML_TYPE_F16, 1); + } + if (!comm->dn16_tensor[i]) { + comm->dn16_tensor[i] = ggml_new_tensor_1d(comm->tctx, GGML_TYPE_F16, 1); + } + comm->up16_tensor[i]->buffer = comm->up16_buffer[i]; + comm->up16_tensor[i]->data = ggml_backend_buffer_get_base(comm->up16_buffer[i]); + comm->dn16_tensor[i]->buffer = comm->dn16_buffer[i]; + comm->dn16_tensor[i]->data = ggml_backend_buffer_get_base(comm->dn16_buffer[i]); + if (!comm->ring_view[i]) { + comm->ring_view[i] = ggml_new_tensor_1d(comm->tctx, GGML_TYPE_F32, 1); + } + } + comm->cap = newcap; + return true; +} + +static bool ggml_backend_vk_comm_allreduce_ring(ggml_backend_vk_comm_context * comm, + ggml_tensor ** tensors) { + const size_t n = comm->backends.size(); + const size_t esz = ggml_type_size(GGML_TYPE_F32); + const size_t xsz = sizeof(uint16_t); + const int64_t nel = ggml_nelements(tensors[0]); + const int64_t cels = (nel + (int64_t) n - 1) / (int64_t) n; + const uint64_t nsteps = 2 * (uint64_t) (n - 1); + const uint64_t nprog = nsteps + 1; + + const uint64_t round = comm->ring_round++; + const size_t slot_off = (size_t) (round & 1) * 2 * comm->cap; + + std::vector compute_val(n), reduce_val(n), prev_reduce(n), up_base(n), pxy_base(n); + for (size_t i = 0; i < n; i++) { + prev_reduce[i] = comm->last_reduce[i]; + compute_val[i] = comm->vkctx[i]->comm_prog_val; + reduce_val[i] = compute_val[i] + nprog; + comm->vkctx[i]->comm_prog_val = reduce_val[i]; + comm->last_reduce[i] = reduce_val[i]; + up_base[i] = comm->up_val[i]; + comm->up_val[i] += nsteps; + if (comm->proxy) { + pxy_base[i] = comm->pxy_val[i]; + comm->pxy_val[i] += nsteps + 1; + } + uint64_t done = comm->device[i]->device.getSemaphoreCounterValue(comm->prog[i]); + if (done >= comm->pool_max_val[i]) { ggml_vk_command_pool_cleanup(comm->device[i], comm->cmd_pool[i]); } + comm->pool_max_val[i] = reduce_val[i]; + uint64_t xdone = comm->device[i]->device.getSemaphoreCounterValue(comm->up[i]); + if (xdone >= comm->xfer_pool_max_val[i]) { ggml_vk_command_pool_cleanup(comm->device[i], comm->cmd_pool_xfer[i]); } + comm->xfer_pool_max_val[i] = up_base[i] + nsteps; + } + + for (size_t i = 0; i < n; i++) { + const size_t nextd = (i + 1) % n; + const size_t prevd = (i + n - 1) % n; + ggml_tensor * rview = comm->ring_view[i]; + char * tbase = (char *) tensors[i]->data; + + vk_context cctx = ggml_vk_create_temporary_context(comm->cmd_pool[i]); + vk_context tctx = ggml_vk_create_temporary_context(comm->cmd_pool_xfer[i]); + + auto chunk_cels = [&](size_t c) -> int64_t { + const int64_t off = (int64_t) c * cels; + return off >= nel ? 0 : std::min(cels, nel - off); + }; + auto set_view = [](ggml_tensor * tt, ggml_backend_buffer_t buf, char * data, int64_t ne0, size_t es) { + tt->ne[0] = ne0; tt->ne[1] = tt->ne[2] = tt->ne[3] = 1; + tt->nb[0] = es; tt->nb[1] = tt->nb[2] = tt->nb[3] = (size_t) ne0 * es; + tt->buffer = buf; tt->data = data; tt->view_offs = 0; + }; + + { + ggml_backend_vk_buffer_context * ubc = (ggml_backend_vk_buffer_context *) comm->up16_buffer[i]->context; + ggml_backend_vk_buffer_context * dbc = (ggml_backend_vk_buffer_context *) comm->dn16_buffer[i]->context; + ggml_tensor * up16t = comm->up16_tensor[i]; + ggml_tensor * dn16t = comm->dn16_tensor[i]; + char * ubase = (char *) ggml_backend_buffer_get_base(comm->up16_buffer[i]); + char * dbase = (char *) ggml_backend_buffer_get_base(comm->dn16_buffer[i]); + + const int64_t cels0 = chunk_cels(i); + ggml_vk_ctx_begin(comm->device[i], cctx); + cctx->s->wait_semaphores.push_back({ comm->prog[i], compute_val[i] }); + cctx->s->wait_semaphores.push_back({ comm->up[i], up_base[i] }); + if (cels0) { + set_view(rview, tensors[i]->buffer, tbase + (size_t) i * cels * esz, cels0, esz); + set_view(up16t, comm->up16_buffer[i], ubase, cels0, xsz); + ggml_vk_cpy(comm->vkctx[i], cctx, rview, up16t); + } + ggml_vk_ctx_end(cctx); + cctx->seqs.back().back().signal_semaphores.push_back({ comm->prog[i], compute_val[i] + 1 }); + + for (uint64_t t = 0; t < nsteps; t++) { + const bool rs = (t < (uint64_t) (n - 1)); + const size_t s = (size_t) (rs ? t : t - (n - 1)); + const size_t c_send = rs ? (i + n - s) % n : (i + n + 1 - s) % n; + const size_t c_recv = rs ? (i + n - s - 1) % n : (i + n - s) % n; + const int64_t scels = chunk_cels(c_send), rcels = chunk_cels(c_recv); + const size_t uoff = (size_t) t * (size_t) cels * xsz; + const size_t hoff = slot_off + uoff; + + ggml_vk_ctx_begin(comm->device[i], tctx); + if (t == 0) { + tctx->s->wait_semaphores.push_back({ comm->prog[i], compute_val[i] + 1 }); + if (round >= 2) { + if (comm->proxy) { + const uint64_t pv = pxy_base[i] + 1; + tctx->s->wait_semaphores.push_back({ comm->pxy[i], pv }); + std::lock_guard lk(comm->bridge_mtx); + comm->bridge_q[i].push_back({ comm->device[nextd]->device, comm->prog[nextd], prev_reduce[nextd], + comm->device[i]->device, comm->pxy[i], pv }); + } else { + tctx->s->wait_semaphores.push_back({ comm->peer_prog[i][nextd], prev_reduce[nextd] }); + } + } + } else { + tctx->s->wait_semaphores.push_back({ comm->prog[i], compute_val[i] + t + 1 }); + } + if (scels) { + ggml_vk_buffer_copy_async(tctx, comm->host_buf[i][i], hoff, ubc->dev_buffer, uoff, (size_t) scels * xsz); + } + ggml_vk_ctx_end(tctx); + tctx->seqs.back().back().signal_semaphores.push_back({ comm->up[i], up_base[i] + t + 1 }); + + ggml_vk_ctx_begin(comm->device[i], cctx); + if (comm->proxy) { + const uint64_t pv = pxy_base[i] + 2 + t; + cctx->s->wait_semaphores.push_back({ comm->pxy[i], pv }); + std::lock_guard lk(comm->bridge_mtx); + comm->bridge_q[i].push_back({ comm->device[prevd]->device, comm->up[prevd], up_base[prevd] + t + 1, + comm->device[i]->device, comm->pxy[i], pv }); + } else { + cctx->s->wait_semaphores.push_back({ comm->peer_up[i][prevd], up_base[prevd] + t + 1 }); + } + if (rcels) { + ggml_vk_buffer_copy_async(cctx, dbc->dev_buffer, 0, comm->host_buf[prevd][i], hoff, (size_t) rcels * xsz); + ggml_vk_sync_buffers(comm->vkctx[i], cctx); + set_view(dn16t, comm->dn16_buffer[i], dbase, rcels, xsz); + set_view(rview, tensors[i]->buffer, tbase + (size_t) c_recv * cels * esz, rcels, esz); + if (rs) { ggml_vk_add(comm->vkctx[i], cctx, rview, dn16t, rview); } + else { ggml_vk_cpy(comm->vkctx[i], cctx, dn16t, rview); } + ggml_vk_sync_buffers(comm->vkctx[i], cctx); + if (t + 1 < nsteps) { + set_view(up16t, comm->up16_buffer[i], ubase + (size_t) (t + 1) * cels * xsz, rcels, xsz); + ggml_vk_cpy(comm->vkctx[i], cctx, rview, up16t); + ggml_vk_sync_buffers(comm->vkctx[i], cctx); + } + } + ggml_vk_ctx_end(cctx); + cctx->seqs.back().back().signal_semaphores.push_back({ comm->prog[i], compute_val[i] + t + 2 }); + } + } + ggml_vk_submit(tctx, {}); + ggml_vk_submit(cctx, {}); + } + return true; +} + +// Recursive halving/doubling all-reduce for a power-of-two device count. Bandwidth-optimal like the ring but +// uses 2*log2(n) cross-device steps instead of 2*(n-1), so its latency scales far better for n>=4. The per-device +// step schedule (peer, send/recv ranges, reduce-vs-copy) is built exactly as in the reference simulation +// scratchpad/tree_sim.py; the GPU plumbing (F16 host staging, F32 accumulate, timeline-semaphore / CPU-proxy +// sync) mirrors ggml_backend_vk_comm_allreduce_ring. Opt-in via GGML_VK_COMM_TREE; the ring stays the default. +// NOTE: only n=2 (and n=2 with forced proxy) is runnable on a 2-GPU box; the multi-step path is exercised by the +// reference simulation and the n=2 plumbing test, not directly at n>=4. +static bool ggml_backend_vk_comm_allreduce_tree(ggml_backend_vk_comm_context * comm, ggml_tensor ** tensors) { + const size_t n = comm->backends.size(); + const size_t esz = ggml_type_size(GGML_TYPE_F32); + const size_t xsz = sizeof(uint16_t); + const int64_t nel = ggml_nelements(tensors[0]); + + int logn = 0; + while (((size_t) 1 << (logn + 1)) <= n) { logn++; } // n is a power of two (checked by caller) + const uint64_t nsteps = 2u * (uint64_t) logn; + const uint64_t nprog = 2u * nsteps; // two compute signals (prep-send, reduce) per step + + // Per-step host stride: the largest single send is ceil(nel/2) F16 elements. A distinct offset per step keeps + // a step's staged data alive until the peer has read it. Double-buffered by round like the ring. + const size_t stride = (size_t) ((nel + 1) / 2) * xsz; + const uint64_t round = comm->ring_round++; + const size_t slot_off = (size_t) (round & 1) * 2 * comm->cap; + + // Build each device's schedule, mirroring tree_sim.py (validated for n=2..8). + struct step_t { size_t peer; int64_t soff, slen, roff, rlen; bool add; }; + std::vector> sched(n); + std::vector> reuse_peers(n); // distinct peers that read host[i] (for cross-call reuse wait) + { + std::vector coff(n, 0), clen(n, nel); + for (size_t dist = n / 2; dist >= 1; dist /= 2) { // reduce-scatter (recursive halving) + for (size_t i = 0; i < n; i++) { + const size_t p = i ^ dist; + const int64_t off = coff[i], len = clen[i]; + const int64_t half = len / 2; + step_t st; st.peer = p; st.add = true; + if (i < p) { st.soff = off + half; st.slen = len - half; st.roff = off; st.rlen = half; coff[i] = off; clen[i] = half; } + else { st.soff = off; st.slen = half; st.roff = off + half; st.rlen = len - half; coff[i] = off + half; clen[i] = len - half; } + sched[i].push_back(st); + } + } + for (size_t dist = 1; dist < n; dist *= 2) { // all-gather (recursive doubling) + std::vector noff(n), nlen(n); + for (size_t i = 0; i < n; i++) { + const size_t p = i ^ dist; + step_t st; st.peer = p; st.add = false; + st.soff = coff[i]; st.slen = clen[i]; + st.roff = coff[p]; st.rlen = clen[p]; + sched[i].push_back(st); + noff[i] = std::min(coff[i], coff[p]); + nlen[i] = std::max(coff[i] + clen[i], coff[p] + clen[p]) - noff[i]; + } + coff = noff; clen = nlen; + } + for (size_t i = 0; i < n; i++) { + for (const auto & st : sched[i]) { + if (std::find(reuse_peers[i].begin(), reuse_peers[i].end(), st.peer) == reuse_peers[i].end()) { + reuse_peers[i].push_back(st.peer); + } + } + } + } + + std::vector compute_val(n), reduce_val(n), prev_reduce(n), up_base(n), pxy_base(n); + for (size_t i = 0; i < n; i++) { + prev_reduce[i] = comm->last_reduce[i]; + compute_val[i] = comm->vkctx[i]->comm_prog_val; + reduce_val[i] = compute_val[i] + nprog; + comm->vkctx[i]->comm_prog_val = reduce_val[i]; + comm->last_reduce[i] = reduce_val[i]; + up_base[i] = comm->up_val[i]; + comm->up_val[i] += nsteps; + if (comm->proxy) { + pxy_base[i] = comm->pxy_val[i]; + comm->pxy_val[i] += nsteps + (uint64_t) reuse_peers[i].size(); + } + uint64_t done = comm->device[i]->device.getSemaphoreCounterValue(comm->prog[i]); + if (done >= comm->pool_max_val[i]) { ggml_vk_command_pool_cleanup(comm->device[i], comm->cmd_pool[i]); } + comm->pool_max_val[i] = reduce_val[i]; + uint64_t xdone = comm->device[i]->device.getSemaphoreCounterValue(comm->up[i]); + if (xdone >= comm->xfer_pool_max_val[i]) { ggml_vk_command_pool_cleanup(comm->device[i], comm->cmd_pool_xfer[i]); } + comm->xfer_pool_max_val[i] = up_base[i] + nsteps; + } + + auto set_view = [](ggml_tensor * tt, ggml_backend_buffer_t buf, char * data, int64_t ne0, size_t es) { + tt->ne[0] = ne0; tt->ne[1] = tt->ne[2] = tt->ne[3] = 1; + tt->nb[0] = es; tt->nb[1] = tt->nb[2] = tt->nb[3] = (size_t) ne0 * es; + tt->buffer = buf; tt->data = data; tt->view_offs = 0; + }; + + for (size_t i = 0; i < n; i++) { + ggml_backend_vk_buffer_context * ubc = (ggml_backend_vk_buffer_context *) comm->up16_buffer[i]->context; + ggml_backend_vk_buffer_context * dbc = (ggml_backend_vk_buffer_context *) comm->dn16_buffer[i]->context; + ggml_tensor * up16t = comm->up16_tensor[i]; + ggml_tensor * dn16t = comm->dn16_tensor[i]; + ggml_tensor * rview = comm->ring_view[i]; + char * ubase = (char *) ggml_backend_buffer_get_base(comm->up16_buffer[i]); + char * dbase = (char *) ggml_backend_buffer_get_base(comm->dn16_buffer[i]); + char * tbase = (char *) tensors[i]->data; + + vk_context cctx = ggml_vk_create_temporary_context(comm->cmd_pool[i]); + vk_context tctx = ggml_vk_create_temporary_context(comm->cmd_pool_xfer[i]); + + for (uint64_t t = 0; t < nsteps; t++) { + const step_t & st = sched[i][t]; + const size_t peer = st.peer; + const size_t hoff = slot_off + (size_t) t * stride; + + // compute-A: copy this step's send range to the F16 staging buffer. + ggml_vk_ctx_begin(comm->device[i], cctx); + cctx->s->wait_semaphores.push_back({ comm->prog[i], compute_val[i] + 2 * t }); + if (st.slen) { + set_view(rview, tensors[i]->buffer, tbase + (size_t) st.soff * esz, st.slen, esz); + set_view(up16t, comm->up16_buffer[i], ubase, st.slen, xsz); + ggml_vk_cpy(comm->vkctx[i], cctx, rview, up16t); + } + ggml_vk_ctx_end(cctx); + cctx->seqs.back().back().signal_semaphores.push_back({ comm->prog[i], compute_val[i] + 2 * t + 1 }); + + // transfer: stage F16 send to this device's host buffer at the step offset. + ggml_vk_ctx_begin(comm->device[i], tctx); + tctx->s->wait_semaphores.push_back({ comm->prog[i], compute_val[i] + 2 * t + 1 }); + if (t == 0 && round >= 2) { + // The round-1-ago call wrote/read the same host slot; ensure those peer reads are done before reuse. + for (size_t pi = 0; pi < reuse_peers[i].size(); pi++) { + const size_t rp = reuse_peers[i][pi]; + if (comm->proxy) { + const uint64_t pv = pxy_base[i] + 1 + pi; + tctx->s->wait_semaphores.push_back({ comm->pxy[i], pv }); + std::lock_guard lk(comm->bridge_mtx); + comm->bridge_q[i].push_back({ comm->device[rp]->device, comm->prog[rp], prev_reduce[rp], + comm->device[i]->device, comm->pxy[i], pv }); + } else { + tctx->s->wait_semaphores.push_back({ comm->peer_prog[i][rp], prev_reduce[rp] }); + } + } + } + if (st.slen) { + ggml_vk_buffer_copy_async(tctx, comm->host_buf[i][i], hoff, ubc->dev_buffer, 0, (size_t) st.slen * xsz); + } + ggml_vk_ctx_end(tctx); + tctx->seqs.back().back().signal_semaphores.push_back({ comm->up[i], up_base[i] + t + 1 }); + + // compute-B: wait our own + the peer's transfer, read the peer's staged data, reduce/copy it in. + ggml_vk_ctx_begin(comm->device[i], cctx); + cctx->s->wait_semaphores.push_back({ comm->up[i], up_base[i] + t + 1 }); + if (comm->proxy) { + const uint64_t pv = pxy_base[i] + (uint64_t) reuse_peers[i].size() + 1 + t; + cctx->s->wait_semaphores.push_back({ comm->pxy[i], pv }); + std::lock_guard lk(comm->bridge_mtx); + comm->bridge_q[i].push_back({ comm->device[peer]->device, comm->up[peer], up_base[peer] + t + 1, + comm->device[i]->device, comm->pxy[i], pv }); + } else { + cctx->s->wait_semaphores.push_back({ comm->peer_up[i][peer], up_base[peer] + t + 1 }); + } + if (st.rlen) { + ggml_vk_buffer_copy_async(cctx, dbc->dev_buffer, 0, comm->host_buf[peer][i], hoff, (size_t) st.rlen * xsz); + ggml_vk_sync_buffers(comm->vkctx[i], cctx); + set_view(dn16t, comm->dn16_buffer[i], dbase, st.rlen, xsz); + set_view(rview, tensors[i]->buffer, tbase + (size_t) st.roff * esz, st.rlen, esz); + if (st.add) { ggml_vk_add(comm->vkctx[i], cctx, rview, dn16t, rview); } + else { ggml_vk_cpy(comm->vkctx[i], cctx, dn16t, rview); } + ggml_vk_sync_buffers(comm->vkctx[i], cctx); + } + ggml_vk_ctx_end(cctx); + cctx->seqs.back().back().signal_semaphores.push_back({ comm->prog[i], compute_val[i] + 2 * t + 2 }); + } + ggml_vk_submit(tctx, {}); + ggml_vk_submit(cctx, {}); + } + return true; +} + +static bool ggml_backend_vk_comm_allreduce_tensor(void * comm_ctx, ggml_tensor ** tensors) { + ggml_backend_vk_comm_context * comm = static_cast(comm_ctx); + const size_t n = comm->backends.size(); + + if (!comm->fast || tensors[0]->type != GGML_TYPE_F32) { + return false; + } + const int64_t ne = ggml_nelements(tensors[0]); + if (ne == 0) { + return true; + } + const size_t nbytes = ggml_nbytes(tensors[0]); + for (size_t i = 0; i < n; i++) { + if (tensors[i]->type != GGML_TYPE_F32 || ggml_nelements(tensors[i]) != ne || !ggml_is_contiguous(tensors[i])) { + return false; + } + } + if (!ggml_backend_vk_comm_ensure(comm, nbytes)) { + return false; + } + + constexpr size_t ring_min = 2u << 20; + if (comm->ring_ok && n >= 2 && nbytes >= ring_min) { + bool all_compute = true; + for (size_t i = 0; i < n; i++) { + all_compute = all_compute && (tensors[i]->flags & GGML_TENSOR_FLAG_COMPUTE); + } + if (all_compute) { + // Opt-in recursive halving/doubling for power-of-two device counts (fewer cross-device steps than the + // ring at n>=4); the ring remains the default and handles non-power-of-two counts. + static const bool use_tree = getenv("GGML_VK_COMM_TREE") != nullptr; + if (use_tree && (n & (n - 1)) == 0) { + return ggml_backend_vk_comm_allreduce_tree(comm, tensors); + } + return ggml_backend_vk_comm_allreduce_ring(comm, tensors); + } + } + + std::vector compute_val(n), gather_val(n), reduce_val(n), prev_reduce(n), pxy_base(n); + for (size_t i = 0; i < n; i++) { + prev_reduce[i] = comm->last_reduce[i]; + compute_val[i] = comm->vkctx[i]->comm_prog_val; + gather_val[i] = compute_val[i] + 1; + reduce_val[i] = compute_val[i] + 2; + comm->vkctx[i]->comm_prog_val = reduce_val[i]; + comm->last_reduce[i] = reduce_val[i]; + if (comm->proxy) { + pxy_base[i] = comm->pxy_val[i]; + comm->pxy_val[i] += 2 * (uint64_t) (n - 1); + } + uint64_t done = comm->device[i]->device.getSemaphoreCounterValue(comm->prog[i]); + if (done < comm->pool_max_val[i] && comm->cmd_pool[i].buffers_in_use() >= 256) { + vk::SemaphoreWaitInfo wi; + wi.semaphoreCount = 1; + wi.pSemaphores = &comm->prog[i]; + wi.pValues = &comm->pool_max_val[i]; + (void) comm->device[i]->device.waitSemaphores(wi, UINT64_MAX); + done = comm->pool_max_val[i]; + } + if (done >= comm->pool_max_val[i]) { + ggml_vk_command_pool_cleanup(comm->device[i], comm->cmd_pool[i]); + } + comm->pool_max_val[i] = reduce_val[i]; + } + + for (size_t i = 0; i < n; i++) { + ggml_tensor * tmp = comm->tmp_tensor[i]; + for (int d = 0; d < GGML_MAX_DIMS; d++) { + tmp->ne[d] = tensors[i]->ne[d]; + } + tmp->nb[0] = ggml_type_size(GGML_TYPE_F32); + for (int d = 1; d < GGML_MAX_DIMS; d++) { + tmp->nb[d] = tmp->nb[d - 1] * tmp->ne[d - 1]; + } + ggml_backend_vk_buffer_context * tbc = (ggml_backend_vk_buffer_context *) comm->tmp_buffer[i]->context; + vk_context c = ggml_vk_create_temporary_context(comm->cmd_pool[i]); + + ggml_vk_ctx_begin(comm->device[i], c); + c->s->wait_semaphores.push_back({ comm->prog[i], compute_val[i] }); + if (comm->proxy) { + c->s->wait_semaphores.push_back({ comm->pxy[i], pxy_base[i] + (uint64_t) (n - 1) }); + } + for (size_t k = 0; k < n; k++) { + if (k != i) { + const size_t p = (k < i) ? k : k - 1; + if (comm->proxy) { + std::lock_guard lk(comm->bridge_mtx); + comm->bridge_q[i].push_back({ comm->device[k]->device, comm->prog[k], prev_reduce[k], + comm->device[i]->device, comm->pxy[i], pxy_base[i] + p + 1 }); + } else { + c->s->wait_semaphores.push_back({ comm->peer_prog[i][k], prev_reduce[k] }); + } + } + } + if (tensors[i]->flags & GGML_TENSOR_FLAG_COMPUTE) { + ggml_backend_vk_buffer_context * bc = (ggml_backend_vk_buffer_context *) tensors[i]->buffer->context; + ggml_vk_buffer_copy_async(c, comm->host_buf[i][i], 0, bc->dev_buffer, + vk_tensor_offset(tensors[i]) + tensors[i]->view_offs, nbytes); + } else { + ggml_vk_buffer_memset_async(c, comm->host_buf[i][i], 0, 0, nbytes); + } + ggml_vk_ctx_end(c); + c->seqs.back().back().signal_semaphores.push_back({ comm->prog[i], gather_val[i] }); + + ggml_vk_ctx_begin(comm->device[i], c); + c->s->wait_semaphores.push_back({ comm->prog[i], gather_val[i] }); + if (comm->proxy) { + c->s->wait_semaphores.push_back({ comm->pxy[i], pxy_base[i] + 2 * (uint64_t) (n - 1) }); + } + for (size_t k = 0; k < n; k++) { + if (k != i) { + const size_t p = (k < i) ? k : k - 1; + if (comm->proxy) { + std::lock_guard lk(comm->bridge_mtx); + comm->bridge_q[i].push_back({ comm->device[k]->device, comm->prog[k], gather_val[k], + comm->device[i]->device, comm->pxy[i], pxy_base[i] + (uint64_t) (n - 1) + p + 1 }); + } else { + c->s->wait_semaphores.push_back({ comm->peer_prog[i][k], gather_val[k] }); + } + } + } + for (size_t k = 0; k < n; k++) { + if (k == i) { + continue; + } + ggml_vk_buffer_copy_async(c, tbc->dev_buffer, 0, comm->host_buf[k][i], 0, nbytes); + ggml_vk_sync_buffers(comm->vkctx[i], c); + ggml_vk_add(comm->vkctx[i], c, tensors[i], tmp, tensors[i]); + ggml_vk_sync_buffers(comm->vkctx[i], c); + } + ggml_vk_ctx_end(c); + c->seqs.back().back().signal_semaphores.push_back({ comm->prog[i], reduce_val[i] }); + + ggml_vk_submit(c, {}); + } + return true; +} + +static void * ggml_backend_vk_reg_get_proc_address(ggml_backend_reg_t reg, const char * name) { + UNUSED(reg); + if (strcmp(name, "ggml_backend_comm_init") == 0) { + return (void *) ggml_backend_vk_comm_init; + } + if (strcmp(name, "ggml_backend_comm_free") == 0) { + return (void *) ggml_backend_vk_comm_free; + } + if (strcmp(name, "ggml_backend_comm_allreduce_tensor") == 0) { + return (void *) ggml_backend_vk_comm_allreduce_tensor; + } + 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() {