From 44dde2dfa2dea7129c12ccbcf3d8e50d6c334984 Mon Sep 17 00:00:00 2001 From: Xun Sun Date: Wed, 17 Dec 2025 15:50:12 +0800 Subject: [PATCH] [EP] Fix the tensorSize of the barrier op (#1222) * Fix the tensorSize of the barrier op * Fix --- mooncake-ep/src/mooncake_backend.cpp | 7 +++++-- mooncake-wheel/tests/test_mooncake_backend_cpu.py | 7 +++++++ 2 files changed, 12 insertions(+), 2 deletions(-) diff --git a/mooncake-ep/src/mooncake_backend.cpp b/mooncake-ep/src/mooncake_backend.cpp index 9fa03df4..aa90e83c 100644 --- a/mooncake-ep/src/mooncake_backend.cpp +++ b/mooncake-ep/src/mooncake_backend.cpp @@ -14,6 +14,7 @@ constexpr const char* SYNC_OP_ERROR_MSG = "Expecting async op but got sync op."; constexpr const char* REDUCE_OP_ERROR_MSG = "Only support SUM."; constexpr const char* SPARSE_ERROR_MSG = "Sparse op not supported."; constexpr const char* REDUCE_DTYPE_ERROR_MSG = "Unsupported reduce dtype: "; +constexpr int kBarrierDummyTensorSize = 1; std::string MooncakeBackend::hostIp_ = "127.0.0.1"; TransferEngine MooncakeBackend::engine_ = TransferEngine(true); @@ -469,8 +470,10 @@ c10::intrusive_ptr MooncakeBackend::barrier( const c10d::BarrierOptions& opts) { TORCH_CHECK(isCpu_, "Barrier is available only for CPU.") return worker_.putTaskCpu( - c10d::OpType::BARRIER, 0, 0, &meta_, [=](void*, size_t, size_t) {}, - [=](void*, size_t, size_t) {}); + // a non-zero tensorSize is required to ensure the worker task for the + // barrier is created + c10d::OpType::BARRIER, kBarrierDummyTensorSize, 0, &meta_, + [=](void*, size_t, size_t) {}, [=](void*, size_t, size_t) {}); } void MooncakeBackend::shutdown() { diff --git a/mooncake-wheel/tests/test_mooncake_backend_cpu.py b/mooncake-wheel/tests/test_mooncake_backend_cpu.py index a806adf0..04e0d028 100644 --- a/mooncake-wheel/tests/test_mooncake_backend_cpu.py +++ b/mooncake-wheel/tests/test_mooncake_backend_cpu.py @@ -26,6 +26,10 @@ def worker(rank, world_size, results, collective): dist.all_gather(gathered, tensor) results[rank] = [t.item() for t in gathered] + elif collective == "barrier": + dist.barrier() + results[rank] = "ok" + else: raise ValueError(f"Unsupported collective: {collective}") @@ -66,6 +70,9 @@ class TestMooncakeBackend(unittest.TestCase): # Expected gather = [0, 1, 2, 3] self._spawn_and_check("all_gather", lambda size: list(range(size))) + def test_barrier(self): + self._spawn_and_check("barrier", lambda size: "ok") + if __name__ == "__main__": unittest.main()