Skip to content

Commit 70f6f4d

Browse files
committed
add busy spin for dispatcher thread when using sm tensor
1 parent 7bfa824 commit 70f6f4d

1 file changed

Lines changed: 44 additions & 1 deletion

File tree

‎ggml/src/ggml-rpc/ggml-rpc.cpp‎

Lines changed: 44 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -383,6 +383,14 @@ static bool send_rpc_cmd(socket_ptr sock, enum rpc_cmd cmd, const void * input,
383383

384384
// RPC client-side implementation
385385

386+
static inline void rpc_cpu_relax() {
387+
#if defined(__aarch64__) && (defined(__clang__) || defined(__GNUC__))
388+
__asm__ volatile("yield" ::: "memory");
389+
#else
390+
std::this_thread::yield();
391+
#endif
392+
}
393+
386394
// Performs HELLO handshake with transport auto-negotiation.
387395
// Advertises local capabilities via conn_caps; if the server responds with
388396
// matching capabilities, the socket is upgraded transparently.
@@ -431,6 +439,16 @@ class message_queue {
431439
return true;
432440
}
433441

442+
bool try_pop(T* out) {
443+
std::unique_lock<std::mutex> lock(mutex);
444+
if (interrupted || queue.empty()) {
445+
return false;
446+
}
447+
*out = queue.front();
448+
queue.pop();
449+
return true;
450+
}
451+
434452
void interrupt() {
435453
std::unique_lock<std::mutex> lock(mutex);
436454
interrupted = true;
@@ -460,6 +478,8 @@ class rpc_dispatcher {
460478
void event_synchronize(ggml_backend_event_t event);
461479
void event_record(ggml_backend_event_t event);
462480
void synchronize();
481+
void busy_spin_acquire();
482+
void busy_spin_release();
463483

464484
void start(const std::string & endpoint);
465485
void work();
@@ -483,6 +503,7 @@ class rpc_dispatcher {
483503
};
484504
rpc_msg_queue queue;
485505
socket_ptr sock;
506+
std::atomic_uint busy_spin_users = 0;
486507
std::atomic_bool running;
487508
std::thread thread;
488509
};
@@ -574,6 +595,15 @@ void rpc_dispatcher::synchronize() {
574595
msg->completion.get_future().wait();
575596
}
576597

598+
void rpc_dispatcher::busy_spin_acquire() {
599+
busy_spin_users.fetch_add(1, std::memory_order_relaxed);
600+
}
601+
602+
void rpc_dispatcher::busy_spin_release() {
603+
const unsigned previous = busy_spin_users.fetch_sub(1, std::memory_order_relaxed);
604+
GGML_ASSERT(previous > 0);
605+
}
606+
577607
void rpc_dispatcher::start(const std::string & endpoint) {
578608
std::string host;
579609
int port;
@@ -599,7 +629,12 @@ void rpc_dispatcher::start(const std::string & endpoint) {
599629
void rpc_dispatcher::work() {
600630
while (running) {
601631
rpc_msg_ptr msg_ptr;
602-
if (!queue.pop(&msg_ptr)) {
632+
if (busy_spin_users.load(std::memory_order_relaxed) != 0) {
633+
if (!queue.try_pop(&msg_ptr)) {
634+
rpc_cpu_relax();
635+
continue;
636+
}
637+
} else if (!queue.pop(&msg_ptr)) {
603638
break;
604639
}
605640
if (msg_ptr->cmd != RPC_CMD_NONE) {
@@ -2792,6 +2827,7 @@ static void ggml_backend_rpc_comm_free(void * comm_ctx_v) {
27922827
auto request = std::make_shared<rpc_msg_comm_free_req>();
27932828
request->device = rank.device;
27942829
rank.dispatcher->send(RPC_CMD_COMM_FREE, request, sizeof(*request));
2830+
rank.dispatcher->busy_spin_release();
27952831
}
27962832
delete comm_ctx;
27972833
}
@@ -2830,6 +2866,10 @@ static void * ggml_backend_rpc_comm_init(ggml_backend_t * backends, size_t n_bac
28302866
return nullptr;
28312867
}
28322868

2869+
for (const auto & rank : ranks) {
2870+
rank.dispatcher->busy_spin_acquire();
2871+
}
2872+
28332873
// Send all init requests before reading any response: rank 0 blocks in accept
28342874
// until rank 1 has connected.
28352875
std::vector<rpc_msg_comm_init_rsp> responses(n_backends);
@@ -2855,6 +2895,9 @@ static void * ggml_backend_rpc_comm_init(ggml_backend_t * backends, size_t n_bac
28552895
}
28562896
}
28572897
if (!ok) {
2898+
for (const auto & rank : ranks) {
2899+
rank.dispatcher->busy_spin_release();
2900+
}
28582901
return nullptr;
28592902
}
28602903
GGML_LOG_INFO("%s: pairwise communicator initialized (%s <-> %s)\n", __func__,

0 commit comments

Comments
 (0)