Merge branch 'pr-25051' into 0cc4m/vulkan-allreduce

This commit is contained in:
Ruben Ortlam
2026-08-25 12:26:42 +02:00
4 changed files with 938 additions and 4 deletions
+3
View File
@@ -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)
+15 -1
View File
@@ -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];
+14
View File
@@ -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));
+906 -3
View File
@@ -67,10 +67,12 @@ typedef struct VkPhysicalDeviceCooperativeMatrixDecodeVectorFeaturesNV {
#include <future>
#include <condition_variable>
#include <thread>
#include <atomic>
#if defined(_MSC_VER)
# define NOMINMAX 1
# include <windows.h>
# include <malloc.h> // _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<ggml_backend_t> backends;
std::vector<ggml_backend_vk_context*> vkctx;
std::vector<vk_device> device;
size_t align = 4096;
size_t cap = 0;
bool fast = false;
std::vector<vk::Semaphore> prog;
std::vector<std::vector<vk::Semaphore>> peer_prog;
std::vector<uint64_t> last_reduce;
std::vector<vk_command_pool> cmd_pool;
std::vector<uint64_t> pool_max_val;
PFN_vkGetSemaphoreFdKHR pGetSemFd = nullptr;
PFN_vkImportSemaphoreFdKHR pImportSemFd = nullptr;
std::vector<void*> host_ptr;
std::vector<std::vector<vk_buffer>> host_buf;
std::vector<ggml_backend_buffer_t> tmp_buffer;
std::vector<ggml_tensor*> tmp_tensor;
ggml_context * tctx = nullptr;
bool ring_ok = false;
std::vector<vk::Semaphore> up;
std::vector<std::vector<vk::Semaphore>> peer_up;
std::vector<uint64_t> up_val;
std::vector<vk_command_pool> cmd_pool_xfer;
std::vector<uint64_t> xfer_pool_max_val;
uint64_t ring_round = 0;
std::vector<ggml_tensor*> ring_view;
std::vector<ggml_backend_buffer_t> up16_buffer;
std::vector<ggml_backend_buffer_t> dn16_buffer;
std::vector<ggml_tensor*> up16_tensor;
std::vector<ggml_tensor*> dn16_tensor;
bool proxy = false;
std::vector<vk::Semaphore> pxy;
std::vector<uint64_t> 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<std::deque<bridge_task>> bridge_q;
std::mutex bridge_mtx;
std::thread proxy_thread;
std::atomic<bool> 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<std::mutex> 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<vk::Semaphore>(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<vk::Semaphore>(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<ggml_backend_vk_comm_context *>(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<uint64_t> 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<std::mutex> 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<std::mutex> 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<std::vector<step_t>> sched(n);
std::vector<std::vector<size_t>> reuse_peers(n); // distinct peers that read host[i] (for cross-call reuse wait)
{
std::vector<int64_t> 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<int64_t> 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<uint64_t> 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<std::mutex> 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<std::mutex> 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<ggml_backend_vk_comm_context *>(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<uint64_t> 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<std::mutex> 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<std::mutex> 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() {