@@ -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+
577607void 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) {
599629void 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