vulkan: when using transfer queue for async copies, sync on event_wait to avoid race (#25229)
This commit is contained in:
@@ -7829,6 +7829,20 @@ static vk_context ggml_vk_get_compute_ctx(ggml_backend_vk_context * ctx) {
|
|||||||
return result;
|
return result;
|
||||||
}
|
}
|
||||||
|
|
||||||
|
static vk_context ggml_vk_get_transfer_ctx(ggml_backend_vk_context * ctx) {
|
||||||
|
vk_context result;
|
||||||
|
if (!ctx->transfer_ctx.expired()) {
|
||||||
|
result = ctx->transfer_ctx.lock();
|
||||||
|
} else {
|
||||||
|
result = ggml_vk_create_context(ctx, ctx->transfer_cmd_pool);
|
||||||
|
|
||||||
|
ctx->transfer_ctx = result;
|
||||||
|
ggml_vk_ctx_begin(ctx->device, result);
|
||||||
|
}
|
||||||
|
|
||||||
|
return result;
|
||||||
|
}
|
||||||
|
|
||||||
// Submit any pending transfer queue work and signal the transfer semaphore.
|
// Submit any pending transfer queue work and signal the transfer semaphore.
|
||||||
// The next compute context created via ggml_vk_get_compute_ctx will wait on this semaphore.
|
// The next compute context created via ggml_vk_get_compute_ctx will wait on this semaphore.
|
||||||
// Returns true if work was submitted.
|
// Returns true if work was submitted.
|
||||||
@@ -15633,13 +15647,7 @@ static void ggml_backend_vk_set_tensor_2d_async(ggml_backend_t backend, ggml_ten
|
|||||||
vk_context cpy_ctx;
|
vk_context cpy_ctx;
|
||||||
|
|
||||||
if (ctx->device->async_use_transfer_queue) {
|
if (ctx->device->async_use_transfer_queue) {
|
||||||
if (ctx->transfer_ctx.expired()) {
|
cpy_ctx = ggml_vk_get_transfer_ctx(ctx);
|
||||||
cpy_ctx = ggml_vk_create_context(ctx, ctx->transfer_cmd_pool);
|
|
||||||
ctx->transfer_ctx = cpy_ctx;
|
|
||||||
ggml_vk_ctx_begin(ctx->device, cpy_ctx);
|
|
||||||
} else {
|
|
||||||
cpy_ctx = ctx->transfer_ctx.lock();
|
|
||||||
}
|
|
||||||
} else {
|
} else {
|
||||||
cpy_ctx = ggml_vk_get_compute_ctx(ctx);
|
cpy_ctx = ggml_vk_get_compute_ctx(ctx);
|
||||||
}
|
}
|
||||||
@@ -15785,13 +15793,7 @@ static bool ggml_backend_vk_cpy_tensor_async(ggml_backend_t backend_src, ggml_ba
|
|||||||
|
|
||||||
vk_context cpy_ctx;
|
vk_context cpy_ctx;
|
||||||
if (ctx->device->async_use_transfer_queue) {
|
if (ctx->device->async_use_transfer_queue) {
|
||||||
if (ctx->transfer_ctx.expired()) {
|
cpy_ctx = ggml_vk_get_transfer_ctx(ctx);
|
||||||
cpy_ctx = ggml_vk_create_context(ctx, ctx->transfer_cmd_pool);
|
|
||||||
ctx->transfer_ctx = cpy_ctx;
|
|
||||||
ggml_vk_ctx_begin(ctx->device, cpy_ctx);
|
|
||||||
} else {
|
|
||||||
cpy_ctx = ctx->transfer_ctx.lock();
|
|
||||||
}
|
|
||||||
} else {
|
} else {
|
||||||
cpy_ctx = ggml_vk_get_compute_ctx(ctx);
|
cpy_ctx = ggml_vk_get_compute_ctx(ctx);
|
||||||
}
|
}
|
||||||
@@ -17123,6 +17125,11 @@ static void ggml_backend_vk_event_wait(ggml_backend_t backend, ggml_backend_even
|
|||||||
if (vkev->has_event) {
|
if (vkev->has_event) {
|
||||||
// Wait for latest event
|
// Wait for latest event
|
||||||
ggml_vk_wait_events(compute_ctx, { vkev->event });
|
ggml_vk_wait_events(compute_ctx, { vkev->event });
|
||||||
|
|
||||||
|
if (ctx->device->async_use_transfer_queue) {
|
||||||
|
vk_context transfer_ctx = ggml_vk_get_transfer_ctx(ctx);
|
||||||
|
transfer_ctx->s->wait_semaphores.push_back(vkev->tl_semaphore);
|
||||||
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user