diff --git a/.gitignore b/.gitignore index 22da42161..a874196b9 100644 --- a/.gitignore +++ b/.gitignore @@ -111,4 +111,4 @@ thirdparty/gdrcopy/ .worktrees/ worktrees/ *.csv -*.egg-info/ \ No newline at end of file +*.egg-info//core diff --git a/ep/bench/buffer.py b/ep/bench/buffer.py index dc46e0fea..0f23c8ca3 100644 --- a/ep/bench/buffer.py +++ b/ep/bench/buffer.py @@ -774,7 +774,7 @@ def get_combine_config(num_ranks: int) -> Config: config_map = { 2: Config(Buffer.num_sms, 10, 256, 6, 128), 4: Config(Buffer.num_sms, 9, 256, 6, 128), - 8: Config(Buffer.num_sms, 4, 256, 6, 128), + 8: Config(Buffer.num_sms, 4, 256, 8, 128), 16: Config(Buffer.num_sms, 4, 288, 12, 512 if Buffer._is_efa() else 128), 24: Config(Buffer.num_sms, 1, 288, 8, 128), 32: Config(Buffer.num_sms, 1, 288, 8, 512 if Buffer._is_efa() else 128), diff --git a/ep/bench/dispatch_loop.py b/ep/bench/dispatch_loop.py new file mode 100644 index 000000000..a92dd6f14 --- /dev/null +++ b/ep/bench/dispatch_loop.py @@ -0,0 +1,90 @@ +"""Sustained dispatch loop for wire-utilization measurement. + +Runs buffer.dispatch() at a fixed config in a tight loop for LOOP_DURATION_S +seconds (torchrun-style, one process per GPU). Bracket externally with CXI +telemetry snapshots to measure true NIC bytes/sec. LOOP_CACHED=1 reuses the +dispatch handle (skips the notify/count exchange per iteration). +""" +import gc, os, time +import torch +import torch.distributed as dist +from buffer import Buffer +from utils import init_dist_under_torchrun +from test_internode import compute_buffer_sizes +from uccl.ep import Config + + +def measure(buffer, group, rank, local_world): + duration = float(os.environ.get("LOOP_DURATION_S", "60")) + num_tokens = int(os.environ.get("LOOP_TOKENS", "4096")) + hidden = int(os.environ.get("LOOP_HIDDEN", "7168")) + num_topk = int(os.environ.get("LOOP_TOPK", "8")) + num_experts = int(os.environ.get("LOOP_EXPERTS", "288")) + nvl_chunk = int(os.environ.get("LOOP_NVL_CHUNK", "32")) + nvl_buf = int(os.environ.get("LOOP_NVL_BUF", "256")) + rdma_chunk = int(os.environ.get("LOOP_RDMA_CHUNK", "64")) + rdma_buf = int(os.environ.get("LOOP_RDMA_BUF", "128")) + cached = os.environ.get("LOOP_CACHED", "0") == "1" + + torch.manual_seed(rank) + x = torch.randn((num_tokens, hidden), dtype=torch.bfloat16, device="cuda") + scores = torch.randn((num_tokens, num_experts), dtype=torch.float32, + device="cuda").abs() + 1.0 + topk_idx = torch.topk(scores, num_topk, dim=-1, largest=True)[1] + topk_weights = torch.ones((num_tokens, num_topk), dtype=torch.float32, + device="cuda") + (num_tokens_per_rank, num_tokens_per_rdma_rank, num_tokens_per_expert, + is_token_in_rank, _) = buffer.get_dispatch_layout(topk_idx, num_experts) + config = Config(24, nvl_chunk, nvl_buf, rdma_chunk, rdma_buf) + args = dict(x=x, num_tokens_per_rank=num_tokens_per_rank, + num_tokens_per_rdma_rank=num_tokens_per_rdma_rank, + is_token_in_rank=is_token_in_rank, + topk_idx=topk_idx, topk_weights=topk_weights, + num_tokens_per_expert=num_tokens_per_expert, config=config) + if cached: + recv = buffer.dispatch(**args) + args = dict(x=x, handle=recv[4], config=config) + for _ in range(5): + buffer.dispatch(**args) + torch.cuda.synchronize() + dist.barrier(group) + if rank == 0: + print(f"[loop] start duration={duration}s tokens={num_tokens} " + f"cached={cached}", flush=True) + t0 = time.time() + iters = 0 + while time.time() - t0 < duration: + buffer.dispatch(**args) + iters += 1 + torch.cuda.synchronize() + elapsed = time.time() - t0 + rdma_tokens = int(num_tokens_per_rdma_rank.sum().item()) - int( + num_tokens_per_rdma_rank[rank // local_world].item()) + bytes_per_iter = rdma_tokens * hidden * 2 + print(f"[loop] rank={rank} iters={iters} elapsed={elapsed:.2f}s " + f"rdma_tokens={rdma_tokens} bytes/iter={bytes_per_iter} " + f"offered_GBps={iters * bytes_per_iter / elapsed / 1e9:.2f}", + flush=True) + + +def main(): + local_rank = int(os.environ["LOCAL_RANK"]) + local_world = int(os.environ.get("LOCAL_WORLD_SIZE", "4")) + rank, world, group = init_dist_under_torchrun(local_rank, local_world) + hidden = int(os.environ.get("LOOP_HIDDEN", "7168")) + nvl_b, rdma_b = compute_buffer_sizes(24, hidden, world) + buffer = Buffer(group, nvl_b, rdma_b, low_latency_mode=False, + num_qps_per_rank=24, explicitly_destroy=True) + measure(buffer, group, rank, local_world) + # All measurement tensors died with measure()s frame; flush deferred + # frees before the buffer tears down the CUDA context. + gc.collect() + torch.cuda.empty_cache() + torch.cuda.synchronize() + dist.barrier(group) + buffer.destroy() + dist.destroy_process_group() + + +if __name__ == "__main__": + main() diff --git a/ep/bench/test_internode.py b/ep/bench/test_internode.py index 4808f2bb8..d23472512 100644 --- a/ep/bench/test_internode.py +++ b/ep/bench/test_internode.py @@ -97,7 +97,7 @@ def test_main( args.num_experts, ) - assert num_experts % num_ranks == 0 and num_local_ranks == 8 + assert num_experts % num_ranks == 0 and num_local_ranks in (4, 8) if local_rank == 0: print( f"[config] num_tokens={num_tokens}, hidden={hidden}, num_topk_groups={num_topk_groups}, num_topk={num_topk}", @@ -394,6 +394,10 @@ def check_data(check_x, recv_gbl_rank_prefix_sum): ), f"{calc_diff(check_topk_weights, ref_topk_weights)}" hash_value += hash_tensor(recv_x) + if getattr(args, "smoke_one", False): + if local_rank == 0: + print("[testing] smoke-one complete", flush=True) + return hash_value # For later tuning dispatch_bf16_rdma_send_bytes = num_rdma_token_sent * hidden * 2 @@ -409,6 +413,24 @@ def check_data(check_x, recv_gbl_rank_prefix_sum): # Tune dispatch performance best_dispatch_results = None fp8_factor = (1 + 4 / 128) / 2 + fixed_dispatch = ( + args.fixed_dispatch_nvl_chunk is not None + or args.fixed_dispatch_rdma_chunk is not None + ) + if fixed_dispatch: + if ( + args.fixed_dispatch_nvl_chunk is None + or args.fixed_dispatch_rdma_chunk is None + ): + raise ValueError( + "--fixed-dispatch-nvl-chunk and --fixed-dispatch-rdma-chunk must be set together" + ) + dispatch_nvl_chunks = (args.fixed_dispatch_nvl_chunk,) + dispatch_rdma_chunks = (args.fixed_dispatch_rdma_chunk,) + else: + dispatch_nvl_chunks = range(4, 45, 4) + dispatch_rdma_chunks = range(4, 129, 8) + for current_x in (x_e4m3, x): best_time, best_results = 1e10, None rdma_send_bytes = ( @@ -421,8 +443,8 @@ def check_data(check_x, recv_gbl_rank_prefix_sum): if isinstance(current_x, tuple) else dispatch_bf16_nvl_recv_bytes ) - for nvl_chunk_size in range(4, 45, 4): - for rdma_chunk_size in range(4, 33, 4): + for nvl_chunk_size in dispatch_nvl_chunks: + for rdma_chunk_size in dispatch_rdma_chunks: config = Config( num_sms, nvl_chunk_size, @@ -431,9 +453,13 @@ def check_data(check_x, recv_gbl_rank_prefix_sum): rdma_buffer_size, ) tune_args = {"x": current_x, "handle": handle, "config": config} + os.environ["UCCL_CXI_PHASE"] = ( + "dispatch_fp8" if isinstance(current_x, tuple) else "dispatch_bf16" + ) t, notify_t = bench_kineto( lambda: buffer.dispatch(**tune_args), ("dispatch", "notify") ) + os.environ.pop("UCCL_CXI_PHASE", None) if t == 0 or notify_t == 0: continue if t < best_time: @@ -486,12 +512,29 @@ def check_data(check_x, recv_gbl_rank_prefix_sum): "num_tokens_per_expert": num_tokens_per_expert, "config": dispatch_config if dispatch_config is not None else config, } + os.environ["UCCL_CXI_PHASE"] = "dispatch_final" recv_x, _, _, _, handle, _ = buffer.dispatch(**dispatch_args) + os.environ.pop("UCCL_CXI_PHASE", None) # Tune combine performance + fixed_combine = ( + args.fixed_combine_nvl_chunk is not None + or args.fixed_combine_rdma_chunk is not None + ) + if fixed_combine: + if args.fixed_combine_nvl_chunk is None or args.fixed_combine_rdma_chunk is None: + raise ValueError( + "--fixed-combine-nvl-chunk and --fixed-combine-rdma-chunk must be set together" + ) + combine_nvl_chunks = (args.fixed_combine_nvl_chunk,) + combine_rdma_chunks = (args.fixed_combine_rdma_chunk,) + else: + combine_nvl_chunks = range(1, 8, 1) + combine_rdma_chunks = range(12 if num_nodes == 2 else 8, 33, 4) + best_time, best_results = 1e10, None - for nvl_chunk_size in range(1, 8, 1): - for rdma_chunk_size in range(12 if num_nodes == 2 else 8, 33, 4): + for nvl_chunk_size in combine_nvl_chunks: + for rdma_chunk_size in combine_rdma_chunks: config = Config( num_sms, nvl_chunk_size, @@ -500,9 +543,11 @@ def check_data(check_x, recv_gbl_rank_prefix_sum): rdma_buffer_size, ) tune_args = {"x": recv_x, "handle": handle, "config": config} + os.environ["UCCL_CXI_PHASE"] = "combine" t, notify_t = bench_kineto( lambda: buffer.combine(**tune_args), ("combine", "notify") ) + os.environ.pop("UCCL_CXI_PHASE", None) if t == 0 or notify_t == 0: continue if local_rank == 0: @@ -567,7 +612,7 @@ def test_loop( explicitly_destroy=True, ) - assert num_local_ranks == 8 and num_ranks > 8 + assert num_local_ranks in (4, 8) and num_ranks > num_local_ranks for seed in range(int(1e9)): if local_rank == 0: @@ -659,6 +704,35 @@ def test_loop( action="store_true", help="whether to test compatibility with low-latency kernels", ) + parser.add_argument( + "--smoke-one", + action="store_true", + help="run only the first dispatch/combine correctness variant", + ) + parser.add_argument( + "--fixed-dispatch-nvl-chunk", + type=int, + default=None, + help="benchmark only this dispatch NVL chunk size", + ) + parser.add_argument( + "--fixed-dispatch-rdma-chunk", + type=int, + default=None, + help="benchmark only this dispatch RDMA chunk size", + ) + parser.add_argument( + "--fixed-combine-nvl-chunk", + type=int, + default=None, + help="benchmark only this combine NVL chunk size", + ) + parser.add_argument( + "--fixed-combine-rdma-chunk", + type=int, + default=None, + help="benchmark only this combine RDMA chunk size", + ) args = parser.parse_args() world_size = int(os.environ["WORLD_SIZE"]) local_world_size = int(os.environ["LOCAL_WORLD_SIZE"]) diff --git a/ep/bench/test_internode_simple.py b/ep/bench/test_internode_simple.py index 1aea8e925..b65c2cd31 100644 --- a/ep/bench/test_internode_simple.py +++ b/ep/bench/test_internode_simple.py @@ -22,9 +22,6 @@ from utils import ( init_dist, detect_ib_hca, - get_cpu_proxies_meta, - initialize_uccl, - destroy_uccl, ) @@ -49,16 +46,12 @@ def test_simple_internode(rank: int, num_ranks: int, group: dist.ProcessGroup): device_index ).multi_processor_count - scratch_nbytes = int(1e9) # 256 MB - scratch = torch.empty( - scratch_nbytes, dtype=torch.uint8, device=f"cuda:{device_index}" - ) - proxies, workers = initialize_uccl(scratch, scratch_nbytes, rank, num_ranks, group) + scratch_nbytes = int(1e9) + buffer = None try: buffer = Buffer( group=group, - rdma_buffer_ptr=scratch.data_ptr(), num_nvl_bytes=0, num_rdma_bytes=int(scratch_nbytes), low_latency_mode=True, @@ -71,20 +64,6 @@ def test_simple_internode(rank: int, num_ranks: int, group: dist.ProcessGroup): if rank == 0: print("[simple-test] ✓ Buffer created successfully", flush=True) - buffer.connect_atomic_buffer(proxies[0]) - - for proxy in proxies: - proxy.calculate_and_set_dispatch_recv_data_offset( - num_tokens, hidden, num_experts - ) - proxy.set_atomic_buffer_ptr(proxies[0].get_atomic_buffer_ptr()) - - if rank == 0: - print( - "[simple-test] ✓ dispatch_recv_data_offset calculated and set by CPU proxy", - flush=True, - ) - cumulative_local_expert_recv_stats = torch.zeros( (num_experts // num_ranks,), dtype=torch.int, device="cuda" ) @@ -129,7 +108,6 @@ def test_simple_internode(rank: int, num_ranks: int, group: dist.ProcessGroup): time.sleep(1) print("[simple-test] ✓ before destroy!", flush=True) - except Exception as e: if rank == 0: import traceback @@ -138,16 +116,14 @@ def test_simple_internode(rank: int, num_ranks: int, group: dist.ProcessGroup): traceback.print_exc() raise - try: - buffer.destroy() - except Exception: - pass - dist.barrier() - print("[simple-test] ✓ Buffer destroyed", flush=True) + if buffer is not None: + try: + buffer.destroy() + except Exception: + pass - destroy_uccl(proxies, workers) - dist.barrier() + print("[simple-test] ✓ Buffer destroyed", flush=True) def test_worker(local_rank: int, num_local_ranks: int): @@ -156,7 +132,6 @@ def test_worker(local_rank: int, num_local_ranks: int): try: test_simple_internode(rank, num_ranks, group) finally: - dist.barrier() dist.destroy_process_group() diff --git a/ep/bench/test_zero_layout_guard.py b/ep/bench/test_zero_layout_guard.py new file mode 100644 index 000000000..396af55db --- /dev/null +++ b/ep/bench/test_zero_layout_guard.py @@ -0,0 +1,114 @@ +""" +Regression smoke for the zero-token internode layout path. + +Run with the normal multi-node launcher, for example: + + JOBID=<2-node-job> NODES=2 NTASKS_PER_NODE=4 WORK= \ + env PYTHONPATH="$WORK/ep/bench:$WORK:$WORK/ep/deep_ep_wrapper" \ + LOCAL_WORLD_SIZE=4 python3 ep/bench/test_zero_layout_guard.py + +The test passes a guarded RDMA-rank count buffer to the raw runtime +get_dispatch_layout(num_tokens=0) entry point. The guard catches regressions +where the zero-token path clears num_ranks ints instead of num_rdma_ranks ints. +""" + +import os + +import torch +import torch.distributed as dist + +from buffer import Buffer +from test_internode import compute_buffer_sizes +from utils import init_dist_under_torchrun + + +def main() -> None: + local_rank = int(os.environ["LOCAL_RANK"]) + local_world_size = int(os.environ.get("LOCAL_WORLD_SIZE", "4")) + buffer = None + + try: + rank, world_size, group = init_dist_under_torchrun( + local_rank, local_world_size + ) + num_nodes = world_size // local_world_size + + hidden = 1024 + num_experts = 64 + num_topk = 4 + num_sms = 24 + num_nvlink_bytes, num_rdma_bytes = compute_buffer_sizes( + num_sms, hidden, world_size + ) + + buffer = Buffer( + group, + num_nvlink_bytes, + num_rdma_bytes, + low_latency_mode=False, + explicitly_destroy=True, + ) + + guard_value = 1234567 + topk_idx = torch.empty((0, num_topk), dtype=torch.int64, device="cuda") + num_tokens_per_rank = torch.full( + (world_size,), guard_value, dtype=torch.int32, device="cuda" + ) + rdma_with_guard = torch.full( + (num_nodes + world_size + 8,), + guard_value, + dtype=torch.int32, + device="cuda", + ) + num_tokens_per_rdma_rank = rdma_with_guard[:num_nodes] + rdma_guard = rdma_with_guard[num_nodes:] + num_tokens_per_expert = torch.full( + (num_experts,), guard_value, dtype=torch.int32, device="cuda" + ) + is_token_in_rank = torch.empty( + (0, world_size), dtype=torch.bool, device="cuda" + ) + + buffer.runtime.get_dispatch_layout( + topk_idx.data_ptr(), + 0, + num_topk, + num_experts, + num_tokens_per_rank.data_ptr(), + num_tokens_per_rdma_rank.data_ptr(), + num_tokens_per_expert.data_ptr(), + is_token_in_rank.data_ptr(), + None, + False, + False, + buffer._ll_compute_stream_ptr(torch.device("cuda", local_rank)), + ) + torch.cuda.synchronize() + + assert torch.equal(num_tokens_per_rank, torch.zeros_like(num_tokens_per_rank)) + assert torch.equal( + num_tokens_per_rdma_rank, torch.zeros_like(num_tokens_per_rdma_rank) + ) + assert torch.equal( + num_tokens_per_expert, torch.zeros_like(num_tokens_per_expert) + ) + if not torch.equal(rdma_guard, torch.full_like(rdma_guard, guard_value)): + raise AssertionError( + f"rank {rank}: rdma guard was clobbered: {rdma_guard.cpu().tolist()}" + ) + + dist.barrier(group) + if rank == 0: + print( + f"[zero-layout-guard] pass world={world_size} num_nodes={num_nodes}", + flush=True, + ) + finally: + if buffer is not None: + buffer.destroy() + if dist.is_initialized(): + dist.destroy_process_group() + + +if __name__ == "__main__": + main() diff --git a/ep/bench/utils.py b/ep/bench/utils.py index 370a7bdd2..461cc1ea7 100644 --- a/ep/bench/utils.py +++ b/ep/bench/utils.py @@ -46,6 +46,8 @@ def hash_tensor(t: torch.Tensor): def init_dist(local_rank: int, num_local_ranks: int): # Set device + torch.cuda.set_device(local_rank) + torch.set_default_device(torch.device(f"cuda:{local_rank}")) # NOTES: you may rewrite this function with your own cluster settings ip = os.getenv("MASTER_ADDR", "127.0.0.1") @@ -555,6 +557,9 @@ def initialize_uccl( use_normal_mode=False, rdma_buffer_is_host_allocated=False, ): + if hasattr(scratch_ptr, "data_ptr"): + scratch_ptr = scratch_ptr.data_ptr() + # Only sweep barriers belonging to OUR mode so we don't stomp on a # coexisting Buffer of the other mode in the same process. The C++ # shm name format is `/uccl_barrier__uid__th`, diff --git a/ep/include/cxi_transport.hpp b/ep/include/cxi_transport.hpp new file mode 100644 index 000000000..c5e3e2260 --- /dev/null +++ b/ep/include/cxi_transport.hpp @@ -0,0 +1,100 @@ +#pragma once + +#ifdef USE_CXI + +#include +#include +#include +#include +#include + +#include +#include +#include +#include +#include + +namespace uccl::cxi { + +struct EndpointInfo { + uint64_t mr_key = 0; + uint64_t size = 0; + uint64_t host_mr_key = 0; + uint64_t host_size = 0; + std::vector ep_name; +}; + +// fi_context2 must stay the first member: the CQ returns &ctx and we map it +// back to the owning WriteContext by address. +struct WriteContext { + fi_context2 ctx{}; + std::vector wr_ids; // ring command ids retired on completion + bool in_use = false; +}; + +class Transport { + public: + Transport() = default; + Transport(Transport const&) = delete; + Transport& operator=(Transport const&) = delete; + ~Transport(); + + void init(int device_index = -1); + void register_cuda_buffer(void* ptr, size_t size); + void register_host_buffer(void* ptr, size_t size); + EndpointInfo local_info() const; + fi_addr_t insert_peer(EndpointInfo const& peer); + void write(fi_addr_t peer, void* local, size_t bytes, uint64_t remote_offset, + uint64_t remote_key, WriteContext* ctx); + // Non-throwing on queue-full: returns 0 on success, -FI_EAGAIN when the TX + // queue is exhausted (caller should poll() and retry). Throws on hard + // errors. Increments outstanding() on success. + int try_write(fi_addr_t peer, void* local, size_t bytes, + uint64_t remote_offset, uint64_t remote_key, + WriteContext* ctx); + void inject_atomic_add64(fi_addr_t peer, int64_t value, + uint64_t remote_offset, uint64_t remote_key); + // Returns 0 on success, -FI_EAGAIN when the TX queue is exhausted. + int try_inject_atomic_add64(fi_addr_t peer, int64_t value, + uint64_t remote_offset, uint64_t remote_key); + void wait(WriteContext* ctx); + bool wait_all(std::vector const& ctxs, + std::atomic const* progress_run = nullptr); + // Non-blocking CQ poll. Fills `done` with up to `max` completed contexts + // and returns the count. Decrements outstanding(). Throws on CQ error. + size_t poll(WriteContext** done, size_t max); + // Block until all outstanding writes complete (or progress_run goes + // false). Completed contexts are marked !in_use but their wr_ids are NOT + // delivered to the caller — only use when retirement bookkeeping has been + // handled elsewhere or does not matter (teardown). + bool drain(std::atomic const* progress_run = nullptr); + size_t outstanding() const { return outstanding_; } + + // Pooled contexts for the async path. Contexts handed out by + // acquire_context() must be returned via release_context() after poll() + // reports them complete. Never mix pooled and caller-owned (stack) + // contexts on the same transport. + WriteContext* acquire_context(); + void release_context(WriteContext* ctx); + + private: + std::vector> ctx_pool_; + std::vector free_ctxs_; + fi_info* info_ = nullptr; + fid_fabric* fabric_ = nullptr; + fid_domain* domain_ = nullptr; + fid_ep* ep_ = nullptr; + fid_cq* cq_ = nullptr; + fid_av* av_ = nullptr; + fid_mr* mr_ = nullptr; + fid_mr* host_mr_ = nullptr; + void* cuda_ptr_ = nullptr; + size_t cuda_size_ = 0; + void* host_ptr_ = nullptr; + size_t host_size_ = 0; + size_t outstanding_ = 0; +}; + +} // namespace uccl::cxi + +#endif // USE_CXI diff --git a/ep/include/ep_config.hpp b/ep/include/ep_config.hpp index 0bd0b4673..2b854db68 100644 --- a/ep/include/ep_config.hpp +++ b/ep/include/ep_config.hpp @@ -276,8 +276,16 @@ struct LowLatencyLayout { // rdma_buffer). If they overflow `kAtomicBufferSize`, kernels will spin // forever waiting for flags. if (atomic_buffer_ptr != nullptr) { +#ifdef USE_CXI + // The CXI barrier reserves the top 4 KiB of the atomic buffer + // (see cxi_barrier_slot_offset in proxy.cpp); EP signaling must not + // reach into it. + EP_HOST_ASSERT(2 * signaling_buffer_bytes_internode_aligned <= + static_cast(kAtomicBufferSize) - 4096); +#else EP_HOST_ASSERT(2 * signaling_buffer_bytes_internode_aligned <= static_cast(kAtomicBufferSize)); +#endif } // Assign pointers diff --git a/ep/include/ep_configs.cuh b/ep/include/ep_configs.cuh index e1618329a..7c1d9f92f 100644 --- a/ep/include/ep_configs.cuh +++ b/ep/include/ep_configs.cuh @@ -1,6 +1,8 @@ #pragma once +#ifndef NUM_MAX_NVL_PEERS #define NUM_MAX_NVL_PEERS 8 +#endif #define NUM_MAX_RDMA_PEERS 20 #define NUM_WORKSPACE_BYTES (32 * 1024 * 1024) #define NUM_MAX_LOCAL_EXPERTS 1024 diff --git a/ep/include/proxy_ctx.hpp b/ep/include/proxy_ctx.hpp index 85f12d4ef..ae244ee3c 100644 --- a/ep/include/proxy_ctx.hpp +++ b/ep/include/proxy_ctx.hpp @@ -1,9 +1,14 @@ #pragma once #include "barrier_local.hpp" +#ifdef USE_CXI +#include "cxi_transport.hpp" +#endif #include "util/gpu_rt.h" #include #include +#include #include +#include #include #include @@ -90,6 +95,32 @@ struct ProxyCtx { uint32_t remote_rkey = 0; uint32_t rkey = 0; +#ifdef USE_CXI + std::unique_ptr cxi_transport; + void* cxi_local_base = nullptr; + uint64_t cxi_local_len = 0; + fi_addr_t cxi_peer_addr = FI_ADDR_UNSPEC; + uint64_t cxi_remote_key = 0; + uint64_t cxi_remote_len = 0; + uint64_t cxi_remote_host_key = 0; + uint64_t cxi_remote_host_len = 0; + + // Async write path bookkeeping. Control atomics to this peer must not + // overtake the data writes posted before them (with FI_CXI_RDZV_THRESHOLD=0 + // every write is rendezvous, so same-TX submission order does NOT order + // atomic placement vs write placement at the target). Each atomic is + // queued with the number of writes posted before it and injected once + // that many writes have completed. + struct CxiPendingAtomic { + int64_t value; + uint64_t remote_offset; + uint64_t threshold; // inject when cxi_writes_completed >= threshold + }; + uint64_t cxi_writes_posted = 0; + uint64_t cxi_writes_completed = 0; + std::deque cxi_pending_atomics; +#endif + #ifdef USE_DMABUF // Chunked MR support — populated when the GPU buffer exceeds the per-MR // size limit (e.g. 2 GiB limit with full IOMMU translation using DMA-BUF). diff --git a/ep/include/rdma.hpp b/ep/include/rdma.hpp index bbbd6bacb..7725076b6 100644 --- a/ep/include/rdma.hpp +++ b/ep/include/rdma.hpp @@ -44,6 +44,15 @@ struct RDMAConnectionInfo { uint32_t data_qp_num[kChannelPerProxy]; // #endif +#ifdef USE_CXI + uint64_t cxi_mr_key = 0; + uint64_t cxi_mr_len = 0; + uint64_t cxi_host_mr_key = 0; + uint64_t cxi_host_mr_len = 0; + uint32_t cxi_ep_name_len = 0; + uint8_t cxi_ep_name[512] = {}; +#endif + #ifdef USE_DMABUF // Chunked MR info — exchanged when the GPU buffer is split across // multiple MRs (with IOMMU DMA-BUF 2 GiB limit). num_mr_chunks == 0 means @@ -406,7 +415,16 @@ void post_rdma_async_batched(ProxyCtx& S, void* buf, size_t num_wrs, std::vector const& wrs_to_post, std::vector const& cmds_to_post, std::vector>& ctxs, - int my_rank, int thread_idx, bool use_normal_mode); + int my_rank, int thread_idx, bool use_normal_mode, + std::unordered_set& acked_wrs); +#ifdef USE_CXI +// Async CXI write path (see rdma.cpp). Fast mode posts writes without +// waiting; completions retire ring slots via cxi_retire_ctx, which the proxy +// loop must call regularly on every connected peer ctx. +bool cxi_fast_mode(); +size_t cxi_retire_ctx(ProxyCtx& ctx, std::unordered_set& acked_wrs, + size_t budget); +#endif void local_process_completions(ProxyCtx& S, std::unordered_set& acked_wrs, int thread_idx, ibv_wc* wc, int ne, diff --git a/ep/include/uccl_proxy.hpp b/ep/include/uccl_proxy.hpp index 89ddc4856..1a09030f2 100644 --- a/ep/include/uccl_proxy.hpp +++ b/ep/include/uccl_proxy.hpp @@ -94,7 +94,7 @@ class UcclProxy { void* gpu_buffer_addr_; std::vector peers_; int local_rank_; - void* atomic_buffer_ptr_; + void* atomic_buffer_ptr_ = nullptr; bool atomic_buffer_is_host_allocated_ = false; // true => cudaFreeHost, false => cudaFree int node_idx_; diff --git a/ep/setup.py b/ep/setup.py index 59f7e61ea..c7de622c1 100644 --- a/ep/setup.py +++ b/ep/setup.py @@ -163,6 +163,7 @@ def run(self): nvcc_dlink = [] extra_link_args = [] use_dmabuf = False + use_cxi = int(os.getenv("USE_CXI", "0")) if torch.version.cuda: # Add CUDA library directory to library_dirs @@ -234,6 +235,19 @@ def run(self): library_dirs.append(Path(efa_home) / "lib") libraries.append("efa") + if use_cxi: + libfabric_home = Path(os.getenv("LIBFABRIC_HOME", "/opt/libfabric")) + nvl_peers = int(os.getenv("UCCL_NUM_MAX_NVL_PEERS", "4")) + print(f"Building with CXI/libfabric transport support ({libfabric_home})") + print(f"Building CXI transport with NUM_MAX_NVL_PEERS={nvl_peers}") + cxx_flags.append("-DUSE_CXI") + nvcc_flags.append("-DUSE_CXI") + cxx_flags.append(f"-DNUM_MAX_NVL_PEERS={nvl_peers}") + nvcc_flags.append(f"-DNUM_MAX_NVL_PEERS={nvl_peers}") + include_dirs.append(libfabric_home / "include") + library_dirs.append(libfabric_home / "lib") + libraries.append("fabric") + # DMA-BUF registration avoids nvidia_peermem/efa_nv_peermem. # Set USE_DMABUF=1 to compile with this path. use_dmabuf = int(os.getenv("USE_DMABUF", "0")) @@ -426,6 +440,7 @@ def run(self): if gpu_name: print(f" > GPU: {gpu_name}") print(f" > EFA Support: {'Yes' if has_efa else 'No'}") + print(f" > CXI Support: {'Yes' if use_cxi else 'No'}") print(f" > GH200 Support: {'Yes' if has_gh200 else 'No'}") print(f" > DMA-BUF Support: {'Yes' if use_dmabuf else 'No'}") print(f" > Device Arch: {device_arch}") diff --git a/ep/src/cxi_transport.cpp b/ep/src/cxi_transport.cpp new file mode 100644 index 000000000..3b1a9733c --- /dev/null +++ b/ep/src/cxi_transport.cpp @@ -0,0 +1,391 @@ +#ifdef USE_CXI + +#include "cxi_transport.hpp" + +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include + +#include +#include +#include +#include + +namespace uccl::cxi { +namespace { + +void check_fi(char const* what, int ret) { + if (ret != 0) { + throw std::runtime_error(std::string(what) + ": " + fi_strerror(-ret)); + } +} + +void check_cuda(char const* what, cudaError_t ret) { + if (ret != cudaSuccess) { + throw std::runtime_error(std::string(what) + ": " + + cudaGetErrorString(ret)); + } +} + +fi_threading threading_hint() { + char const* value = std::getenv("UCCL_CXI_THREADING"); + if (!value || std::strcmp(value, "endpoint") == 0) return FI_THREAD_ENDPOINT; + if (std::strcmp(value, "completion") == 0) return FI_THREAD_COMPLETION; + if (std::strcmp(value, "domain") == 0) return FI_THREAD_DOMAIN; + if (std::strcmp(value, "fid") == 0) return FI_THREAD_FID; + if (std::strcmp(value, "safe") == 0) return FI_THREAD_SAFE; + if (std::strcmp(value, "unspec") == 0) return FI_THREAD_UNSPEC; + throw std::runtime_error(std::string("Invalid UCCL_CXI_THREADING=") + value); +} + +} // namespace + +Transport::~Transport() { + if (host_mr_) fi_close(&host_mr_->fid); + if (mr_) fi_close(&mr_->fid); + if (av_) fi_close(&av_->fid); + if (cq_) fi_close(&cq_->fid); + if (ep_) fi_close(&ep_->fid); + if (domain_) fi_close(&domain_->fid); + if (fabric_) fi_close(&fabric_->fid); + if (info_) fi_freeinfo(info_); +} + +void Transport::init(int device_index) { + fi_info* hints = fi_allocinfo(); + if (!hints) throw std::runtime_error("fi_allocinfo failed"); + + hints->fabric_attr->prov_name = strdup("cxi"); + if (device_index >= 0) { + std::ostringstream domain_name; + domain_name << "cxi" << device_index; + hints->domain_attr->name = strdup(domain_name.str().c_str()); + } + hints->ep_attr->type = FI_EP_RDM; + hints->caps = FI_TAGGED | FI_MSG | FI_HMEM | FI_RMA | FI_READ | FI_WRITE | + FI_ATOMIC | FI_REMOTE_WRITE | FI_DIRECTED_RECV | + FI_LOCAL_COMM | FI_REMOTE_COMM; + hints->mode = FI_CONTEXT | FI_CONTEXT2; + hints->domain_attr->threading = threading_hint(); + hints->domain_attr->control_progress = FI_PROGRESS_UNSPEC; + hints->domain_attr->data_progress = FI_PROGRESS_UNSPEC; + hints->domain_attr->mr_mode = FI_MR_LOCAL | FI_MR_HMEM | FI_MR_ENDPOINT | + FI_MR_VIRT_ADDR | FI_MR_ALLOCATED | + FI_MR_PROV_KEY; + hints->domain_attr->mr_key_size = 2; + hints->tx_attr->msg_order = FI_ORDER_SAS; + hints->rx_attr->msg_order = FI_ORDER_SAS; + + try { + check_fi("fi_getinfo(cxi)", fi_getinfo(FI_VERSION(1, 18), nullptr, nullptr, + 0, hints, &info_)); + { + char const* dc = std::getenv("UCCL_CXI_DELIVERY_COMPLETE"); + if (dc && dc[0] == (char)0x31) { + info_->tx_attr->op_flags |= FI_DELIVERY_COMPLETE; + fprintf(stderr, "[CXI] FI_DELIVERY_COMPLETE enabled on TX\n"); + } + } + check_fi("fi_fabric", fi_fabric(info_->fabric_attr, &fabric_, nullptr)); + check_fi("fi_domain", fi_domain(fabric_, info_, &domain_, nullptr)); + check_fi("fi_endpoint", fi_endpoint(domain_, info_, &ep_, nullptr)); + + fi_cq_attr cq_attr{}; + cq_attr.format = FI_CQ_FORMAT_CONTEXT; + cq_attr.size = 4096; + check_fi("fi_cq_open", fi_cq_open(domain_, &cq_attr, &cq_, nullptr)); + check_fi("fi_ep_bind(cq)", + fi_ep_bind(ep_, &cq_->fid, FI_TRANSMIT | FI_RECV)); + + fi_av_attr av_attr{}; + av_attr.type = FI_AV_TABLE; + check_fi("fi_av_open", fi_av_open(domain_, &av_attr, &av_, nullptr)); + check_fi("fi_ep_bind(av)", fi_ep_bind(ep_, &av_->fid, 0)); + +#ifdef FI_OPT_CUDA_API_PERMITTED + bool cuda_api_permitted = false; + check_fi("fi_setopt(FI_OPT_CUDA_API_PERMITTED)", + fi_setopt(&ep_->fid, FI_OPT_ENDPOINT, + FI_OPT_CUDA_API_PERMITTED, &cuda_api_permitted, + sizeof(cuda_api_permitted))); +#endif + + check_fi("fi_enable", fi_enable(ep_)); + } catch (...) { + fi_freeinfo(hints); + throw; + } + fi_freeinfo(hints); +} + +void Transport::register_cuda_buffer(void* ptr, size_t size) { + if (!domain_ || !ep_) throw std::runtime_error("CXI transport not initialized"); + cuda_ptr_ = ptr; + cuda_size_ = size; + + cudaPointerAttributes attrs{}; + check_cuda("cudaPointerGetAttributes", cudaPointerGetAttributes(&attrs, ptr)); + + iovec iov{}; + iov.iov_base = ptr; + iov.iov_len = size; + + fi_mr_attr mr_attr{}; + mr_attr.mr_iov = &iov; + mr_attr.iov_count = 1; + mr_attr.access = FI_SEND | FI_RECV | FI_READ | FI_WRITE | + FI_REMOTE_WRITE | FI_REMOTE_READ; + mr_attr.iface = FI_HMEM_CUDA; + mr_attr.device.cuda = attrs.device; + + check_fi("fi_mr_regattr(cuda)", fi_mr_regattr(domain_, &mr_attr, 0, &mr_)); + if (info_->domain_attr->mr_mode & FI_MR_ENDPOINT) { + check_fi("fi_mr_bind(ep)", fi_mr_bind(mr_, &ep_->fid, 0)); + check_fi("fi_mr_enable", fi_mr_enable(mr_)); + } +} + +void Transport::register_host_buffer(void* ptr, size_t size) { + if (!domain_ || !ep_) throw std::runtime_error("CXI transport not initialized"); + host_ptr_ = ptr; + host_size_ = size; + + cudaPointerAttributes attrs{}; + auto const cuda_status = cudaPointerGetAttributes(&attrs, ptr); + cudaGetLastError(); + if (cuda_status == cudaSuccess && + (attrs.type == cudaMemoryTypeDevice || + attrs.type == cudaMemoryTypeManaged)) { + iovec iov{}; + iov.iov_base = ptr; + iov.iov_len = size; + + fi_mr_attr mr_attr{}; + mr_attr.mr_iov = &iov; + mr_attr.iov_count = 1; + mr_attr.access = FI_SEND | FI_RECV | FI_READ | FI_WRITE | + FI_REMOTE_WRITE | FI_REMOTE_READ; + mr_attr.iface = FI_HMEM_CUDA; + mr_attr.device.cuda = attrs.device; + + check_fi("fi_mr_regattr(control cuda)", + fi_mr_regattr(domain_, &mr_attr, 0, &host_mr_)); + } else { + check_fi("fi_mr_reg(host)", + fi_mr_reg(domain_, ptr, size, + FI_SEND | FI_RECV | FI_READ | FI_WRITE | + FI_REMOTE_WRITE | FI_REMOTE_READ, + 0, 0, 0, &host_mr_, nullptr)); + } + if (info_->domain_attr->mr_mode & FI_MR_ENDPOINT) { + check_fi("fi_mr_bind(host ep)", fi_mr_bind(host_mr_, &ep_->fid, 0)); + check_fi("fi_mr_enable(host)", fi_mr_enable(host_mr_)); + } +} + +EndpointInfo Transport::local_info() const { + if (!ep_ || !mr_) throw std::runtime_error("CXI endpoint or MR missing"); + EndpointInfo out; + out.mr_key = fi_mr_key(mr_); + out.size = cuda_size_; + if (host_mr_) { + out.host_mr_key = fi_mr_key(host_mr_); + out.host_size = host_size_; + } + out.ep_name.resize(512); + size_t len = out.ep_name.size(); + check_fi("fi_getname", fi_getname(&ep_->fid, out.ep_name.data(), &len)); + out.ep_name.resize(len); + return out; +} + +fi_addr_t Transport::insert_peer(EndpointInfo const& peer) { + if (!av_) throw std::runtime_error("CXI AV missing"); + fi_addr_t addr = FI_ADDR_UNSPEC; + int ret = fi_av_insert(av_, peer.ep_name.data(), 1, &addr, 0, nullptr); + if (ret != 1) check_fi("fi_av_insert", ret < 0 ? ret : -FI_EINVAL); + return addr; +} + +namespace { +[[noreturn]] void throw_cq_error(fid_cq* cq) { + fi_cq_err_entry err{}; + fi_cq_readerr(cq, &err, 0); + char buf[256]; + char const* msg = + fi_cq_strerror(cq, err.prov_errno, err.err_data, buf, sizeof(buf)); + throw std::runtime_error(std::string("CXI CQ error: ") + + (msg ? msg : fi_strerror(err.err))); +} +} // namespace + +// Consume up to `max` completions, marking contexts done. Returns number +// consumed; fills `done` (may be null when the caller doesn't need them). +size_t Transport::poll(WriteContext** done, size_t max) { + if (!cq_) throw std::runtime_error("CXI CQ missing"); + if (outstanding_ == 0 || max == 0) return 0; + size_t n = 0; + while (n < max) { + fi_cq_entry entries[16]{}; + size_t const want = std::min(16, max - n); + ssize_t rc = fi_cq_read(cq_, entries, want); + if (rc > 0) { + for (ssize_t i = 0; i < rc; ++i) { + auto* c = reinterpret_cast(entries[i].op_context); + c->in_use = false; + if (outstanding_ > 0) --outstanding_; + if (done) done[n] = c; + ++n; + } + continue; + } + if (rc == -FI_EAGAIN) break; + if (rc == -FI_EAVAIL) throw_cq_error(cq_); + check_fi("fi_cq_read", static_cast(rc)); + } + return n; +} + +int Transport::try_write(fi_addr_t peer, void* local, size_t bytes, + uint64_t remote_offset, uint64_t remote_key, + WriteContext* ctx) { + if (!ep_ || !mr_) throw std::runtime_error("CXI endpoint or MR missing"); + if (!ctx) throw std::runtime_error("CXI write context is null"); + uintptr_t const local_addr = reinterpret_cast(local); + uintptr_t const cuda_base = reinterpret_cast(cuda_ptr_); + uintptr_t const host_base = reinterpret_cast(host_ptr_); + fid_mr* local_mr = nullptr; + if (local_addr >= cuda_base && local_addr + bytes <= cuda_base + cuda_size_) { + local_mr = mr_; + } else if (host_mr_ && local_addr >= host_base && + local_addr + bytes <= host_base + host_size_) { + local_mr = host_mr_; + } else { + throw std::runtime_error("CXI write local range is outside registered MRs"); + } + + // CXI advertises FI_MR_VIRT_ADDR, but the validated CUDA write path on + // Isambard requires offset-zero RMA addressing with provider keys. + ssize_t rc = fi_write(ep_, local, bytes, fi_mr_desc(local_mr), peer, + remote_offset, remote_key, &ctx->ctx); + if (rc == -FI_EAGAIN) return -FI_EAGAIN; + check_fi("fi_write(cuda)", static_cast(rc)); + ctx->in_use = true; + ++outstanding_; + return 0; +} + +void Transport::write(fi_addr_t peer, void* local, size_t bytes, + uint64_t remote_offset, uint64_t remote_key, + WriteContext* ctx) { + uint32_t spins = 0; + for (;;) { + int rc = try_write(peer, local, bytes, remote_offset, remote_key, ctx); + if (rc == 0) return; + // TX queue full: make progress by consuming completions, then retry. + // NOTE: completions consumed here are not reported to any caller, so + // this blocking variant is only safe on transports whose retirement + // bookkeeping happens via wait_all() (the sync path). + poll(nullptr, 16); + if ((++spins & 0x3ff) == 0) sched_yield(); + } +} + +int Transport::try_inject_atomic_add64(fi_addr_t peer, int64_t value, + uint64_t remote_offset, + uint64_t remote_key) { + if (!ep_) throw std::runtime_error("CXI endpoint missing"); + ssize_t rc = fi_inject_atomic(ep_, &value, 1, peer, remote_offset, + remote_key, FI_INT64, FI_SUM); + if (rc == -FI_EAGAIN) return -FI_EAGAIN; + check_fi("fi_inject_atomic(add64)", static_cast(rc)); + return 0; +} + +void Transport::inject_atomic_add64(fi_addr_t peer, int64_t value, + uint64_t remote_offset, + uint64_t remote_key) { + uint32_t spins = 0; + while (try_inject_atomic_add64(peer, value, remote_offset, remote_key) == + -FI_EAGAIN) { + poll(nullptr, 16); + if ((++spins & 0x3ff) == 0) sched_yield(); + } +} + +void Transport::wait(WriteContext* ctx) { + if (!cq_ || !ctx) throw std::runtime_error("CXI CQ or context missing"); + uint32_t spins = 0; + while (ctx->in_use) { + if (poll(nullptr, 16) == 0) { + if ((++spins & 0x3ff) == 0) sched_yield(); + } + } +} + +bool Transport::wait_all(std::vector const& ctxs, + std::atomic const* progress_run) { + if (!cq_) throw std::runtime_error("CXI CQ missing"); + if (ctxs.empty()) return true; + + auto all_done = [&]() { + for (WriteContext* ctx : ctxs) { + if (!ctx) throw std::runtime_error("CXI write context is null"); + if (ctx->in_use) return false; + } + return true; + }; + + uint32_t spins = 0; + while (!all_done()) { + if (poll(nullptr, 16) == 0) { + if (progress_run && !progress_run->load(std::memory_order_acquire)) { + return false; + } + if ((++spins & 0x3ff) == 0) sched_yield(); + } + } + return true; +} + +WriteContext* Transport::acquire_context() { + if (!free_ctxs_.empty()) { + WriteContext* c = free_ctxs_.back(); + free_ctxs_.pop_back(); + c->wr_ids.clear(); + return c; + } + ctx_pool_.push_back(std::make_unique()); + return ctx_pool_.back().get(); +} + +void Transport::release_context(WriteContext* ctx) { + if (!ctx) return; + ctx->wr_ids.clear(); + ctx->in_use = false; + free_ctxs_.push_back(ctx); +} + +bool Transport::drain(std::atomic const* progress_run) { + uint32_t spins = 0; + while (outstanding_ > 0) { + if (poll(nullptr, 16) == 0) { + if (progress_run && !progress_run->load(std::memory_order_acquire)) { + return false; + } + if ((++spins & 0x3ff) == 0) sched_yield(); + } + } + return true; +} + +} // namespace uccl::cxi + +#endif // USE_CXI diff --git a/ep/src/internode.cu b/ep/src/internode.cu index e37968036..0c2fa5710 100644 --- a/ep/src/internode.cu +++ b/ep/src/internode.cu @@ -13,10 +13,20 @@ namespace uccl { namespace internode { +__device__ __forceinline__ uint64_t pack_token_in_nvl_ranks( + bool const* is_token_in_nvl_ranks) { + uint64_t packed = 0; +#pragma unroll + for (int i = 0; i < NUM_MAX_NVL_PEERS; ++i) { + packed |= static_cast(is_token_in_nvl_ranks[i]) << (i * 8); + } + return packed; +} + struct SourceMeta { int src_rdma_rank, is_token_in_nvl_rank_bits; - EP_STATIC_ASSERT(NUM_MAX_NVL_PEERS == 8, + EP_STATIC_ASSERT(NUM_MAX_NVL_PEERS <= 8, "Invalid number of maximum NVL peers"); __forceinline__ SourceMeta() = default; @@ -366,17 +376,16 @@ __global__ void notify_dispatch( int total_count = 0, per_nvl_rank_count[NUM_MAX_NVL_PEERS] = {0}; for (int64_t i = token_start_idx + lane_id; i < token_end_idx; i += WARP_SIZE) { - EP_STATIC_ASSERT(NUM_MAX_NVL_PEERS * sizeof(bool) == sizeof(uint64_t), - "Invalid number of NVL peers"); - auto is_token_in_rank_uint64 = *reinterpret_cast( - is_token_in_rank + i * num_ranks + - dst_rdma_rank * NUM_MAX_NVL_PEERS); - auto is_token_in_rank_values = - reinterpret_cast(&is_token_in_rank_uint64); + bool has_token_in_rdma_rank = false; #pragma unroll - for (int j = 0; j < NUM_MAX_NVL_PEERS; ++j) - per_nvl_rank_count[j] += is_token_in_rank_values[j]; - total_count += (is_token_in_rank_uint64 != 0); + for (int j = 0; j < NUM_MAX_NVL_PEERS; ++j) { + bool const is_in_nvl_rank = + is_token_in_rank[i * num_ranks + + dst_rdma_rank * NUM_MAX_NVL_PEERS + j]; + per_nvl_rank_count[j] += is_in_nvl_rank; + has_token_in_rdma_rank |= is_in_nvl_rank; + } + total_count += has_token_in_rdma_rank; } // Warp reduce @@ -563,7 +572,7 @@ __global__ void __launch_bounds__( EP_DEVICE_ASSERT(num_topk <= WARP_SIZE); // RDMA symmetric layout - EP_STATIC_ASSERT(NUM_MAX_NVL_PEERS * sizeof(bool) == sizeof(uint64_t), + EP_STATIC_ASSERT(NUM_MAX_NVL_PEERS * sizeof(bool) <= sizeof(uint64_t), "Invalid number of NVL peers"); auto hidden_bytes = hidden_int4 * sizeof(int4); auto scale_bytes = num_scales * sizeof(float); @@ -757,9 +766,9 @@ __global__ void __launch_bounds__( // Read RDMA rank existence uint64_t is_token_in_rank_uint64 = 0; if (lane_id < kNumRDMARanks) { - is_token_in_rank_uint64 = *(reinterpret_cast( + is_token_in_rank_uint64 = pack_token_in_nvl_ranks( is_token_in_rank + token_idx * num_ranks + - lane_id * NUM_MAX_NVL_PEERS)); + lane_id * NUM_MAX_NVL_PEERS); } // Acquire sequential lock @@ -806,9 +815,9 @@ __global__ void __launch_bounds__( // Read RDMA rank existence uint64_t is_token_in_rank_uint64 = 0; if (lane_id < kNumRDMARanks) { - is_token_in_rank_uint64 = __ldg(reinterpret_cast( + is_token_in_rank_uint64 = pack_token_in_nvl_ranks( is_token_in_rank + token_idx * num_ranks + - lane_id * NUM_MAX_NVL_PEERS)); + lane_id * NUM_MAX_NVL_PEERS); global_rdma_tail_idx += (is_token_in_rank_uint64 != 0); } @@ -1092,10 +1101,6 @@ __global__ void __launch_bounds__( translate_dst_rdma_rank(dst_rdma_rank, nvl_rank), channel_id, // NOTE(MaoZiming): use channel_id for rb. lane_id, 0, d2h_channel_addrs, num_d2h_channel_addrs, false, -1, - // NOTE(MaoZiming): for AMD GPUs, we directly send a subsequent RDMA - // to update the tail. For other GPUs and EFA NICs, we use the - // CPU-emulated atomics, allow us to piggyback the atomic operation - // with the RDMA send. #ifndef EFA 0, 0 #else diff --git a/ep/src/internode_ll.cu b/ep/src/internode_ll.cu index 1cb908b89..765f6230c 100644 --- a/ep/src/internode_ll.cu +++ b/ep/src/internode_ll.cu @@ -506,48 +506,52 @@ LOW_LATENCY_DISPATCH_RECV: EP_DEVICE_ASSERT(num_warps_per_group > 1 and num_warp_groups < 15); #endif if (sub_warp_id == 1 and lane_id == 0) { + bool const recv_via_local_or_ipc = + src_rank == rank || + ((src_rank / max_nvl_peers == rank / max_nvl_peers) && + ipc_rdma_base_ptrs != nullptr && + ipc_rdma_base_ptrs[rank % max_nvl_peers] != nullptr && + ipc_rdma_base_ptrs[src_rank % max_nvl_peers] != nullptr); auto start_time = clock64(); - while ((src_rank / max_nvl_peers == rank / max_nvl_peers) && + while (recv_via_local_or_ipc && (num_recv_tokens_ipc = ld_acquire_sys_global( rdma_recv_count + local_expert_idx * num_ranks + src_rank)) == - 0) + 0) { #if defined(__HIP_PLATFORM_AMD__) || defined(__HIPCC__) __builtin_amdgcn_s_sleep(1); -#else - ; #endif + } - while ((src_rank / max_nvl_peers != rank / max_nvl_peers) && + while (!recv_via_local_or_ipc && (num_recv_tokens_internode = static_cast(ld_acquire_sys_global( reinterpret_cast( rdma_recv_count_internode + - local_expert_idx * num_ranks + src_rank)))) == 0) + local_expert_idx * num_ranks + src_rank)))) == 0) { #if defined(__HIP_PLATFORM_AMD__) || defined(__HIPCC__) __builtin_amdgcn_s_sleep(1); -#else - ; #endif + } - if (src_rank / max_nvl_peers == rank / max_nvl_peers) { + if (recv_via_local_or_ipc) { if (ld_acquire_sys_global( reinterpret_cast(rdma_recv_count_internode + local_expert_idx * num_ranks + src_rank)) != 0) { printf( - "Same node but rdma_recv_count_internode is not zero! src_rank: " - "%d, rank: %d, max_nvl_peers: %d\n", + "Local/IPC receive but rdma_recv_count_internode is not zero! " + "src_rank: %d, rank: %d, max_nvl_peers: %d\n", src_rank, rank, max_nvl_peers); EP_DEVICE_ASSERT(false); } } - if (src_rank / max_nvl_peers != rank / max_nvl_peers) { + if (!recv_via_local_or_ipc) { if (ld_acquire_sys_global( rdma_recv_count + local_expert_idx * num_ranks + src_rank) != 0) { printf( - "Different node but rdma_recv_count is not zero! src_rank: %d, " - "rank: %d, max_nvl_peers: %d\n", + "Internode receive but rdma_recv_count is not zero! src_rank: " + "%d, rank: %d, max_nvl_peers: %d\n", src_rank, rank, max_nvl_peers); EP_DEVICE_ASSERT(false); } @@ -1091,43 +1095,47 @@ LOW_LATENCY_COMBINE_RECV: EP_DEVICE_ASSERT(num_warps_per_group > 1); if (sub_warp_id == 0 and lane_id == 0) { auto const src_rank = responsible_expert_idx / num_local_experts; + bool const recv_via_local_or_ipc = + src_rank == rank || + ((src_rank / max_nvl_peers == rank / max_nvl_peers) && + ipc_rdma_base_ptrs != nullptr && + ipc_rdma_base_ptrs[rank % max_nvl_peers] != nullptr && + ipc_rdma_base_ptrs[src_rank % max_nvl_peers] != nullptr); auto start_time = clock64(); - while ((src_rank / max_nvl_peers == rank / max_nvl_peers) && + while (recv_via_local_or_ipc && ld_acquire_sys_global( - rdma_recv_flag + responsible_expert_idx) == 0) + rdma_recv_flag + responsible_expert_idx) == 0) { #if defined(__HIP_PLATFORM_AMD__) || defined(__HIPCC__) __builtin_amdgcn_s_sleep(1); -#else - ; #endif + } - while ((src_rank / max_nvl_peers != rank / max_nvl_peers) && + while (!recv_via_local_or_ipc && ld_acquire_sys_global( reinterpret_cast( - rdma_recv_flag_internode + responsible_expert_idx)) == 0) + rdma_recv_flag_internode + responsible_expert_idx)) == 0) { #if defined(__HIP_PLATFORM_AMD__) || defined(__HIPCC__) __builtin_amdgcn_s_sleep(1); -#else - ; #endif + } - if (src_rank / max_nvl_peers == rank / max_nvl_peers) { + if (recv_via_local_or_ipc) { if (ld_acquire_sys_global( reinterpret_cast( rdma_recv_flag_internode + responsible_expert_idx)) != 0) { printf( - "Same node but rdma_recv_flag_internode is not zero! src_rank: " - "%d, rank: %d, max_nvl_peers: %d\n", + "Local/IPC receive but rdma_recv_flag_internode is not zero! " + "src_rank: %d, rank: %d, max_nvl_peers: %d\n", src_rank, rank, max_nvl_peers); EP_DEVICE_ASSERT(false); } } - if (src_rank / max_nvl_peers != rank / max_nvl_peers) { + if (!recv_via_local_or_ipc) { if (ld_acquire_sys_global( rdma_recv_flag + responsible_expert_idx) != 0) { printf( - "Different node but rdma_recv_flag is not zero! src_rank: %d, " - "rank: %d, max_nvl_peers: %d\n", + "Internode receive but rdma_recv_flag is not zero! src_rank: " + "%d, rank: %d, max_nvl_peers: %d\n", src_rank, rank, max_nvl_peers); EP_DEVICE_ASSERT(false); } diff --git a/ep/src/proxy.cpp b/ep/src/proxy.cpp index b99311b2d..448cb2657 100644 --- a/ep/src/proxy.cpp +++ b/ep/src/proxy.cpp @@ -5,6 +5,7 @@ #include "rdma.hpp" #include "util/util.h" #include // for htonl, ntohl +#include #include #include #include @@ -85,6 +86,39 @@ void unmap_local_barrier_shm(std::string const& name, LocalBarrier* lb, } #endif +static int normal_mode_rank_stride(int num_ranks, int num_nodes) { + if (num_nodes > 0 && num_ranks % num_nodes == 0) { + return std::max(1, num_ranks / num_nodes); + } + return MAX_NUM_GPUS; +} + +static bool skip_network_peer_on_same_ip(bool same_ip, bool use_normal_mode) { +#ifdef USE_CXI + // Low-latency CXI may use host-pinned RDMA buffers, so CUDA IPC can be + // unavailable even for same-node peers. Keep those peer contexts connected. + return same_ip && use_normal_mode; +#else + (void)use_normal_mode; + return same_ip; +#endif +} + +#ifdef USE_CXI +static constexpr size_t kCxiBarrierBytes = 4096; + +static size_t cxi_barrier_slot_offset(int slot) { + return kAtomicBufferSize - kCxiBarrierBytes + + static_cast(slot) * sizeof(int64_t); +} + +static std::atomic* cxi_barrier_slots(void* atomic_buffer_ptr) { + return reinterpret_cast*>( + static_cast(atomic_buffer_ptr) + kAtomicBufferSize - + kCxiBarrierBytes); +} +#endif + Proxy::Proxy(Config const& cfg) : cfg_(cfg) { // Initialize state tracking for each ring buffer listen_port_ = uccl::create_listen_socket(&listen_fd_); @@ -109,6 +143,21 @@ uint64_t Proxy::completed_wr() const { return completion_count_; } void Proxy::pin_thread_to_cpu_wrapper() { if (cfg_.pin_thread) { +#ifdef USE_CXI + int const cxi_numa_node = std::max(0, cfg_.local_rank % 4); + pin_thread_unique(cxi_numa_node, cfg_.local_rank, cfg_.thread_idx, + kNumProxyThs); + int cpu = sched_getcpu(); + if (cpu == -1) { + perror("sched_getcpu"); + } else { + printf( + "Local CPU thread pinned to NUMA node %d, core %d, thread_idx: %d, " + "local_rank: %d, mode: %s\n", + cxi_numa_node, cpu, cfg_.thread_idx, cfg_.local_rank, + cfg_.use_normal_mode ? "high_throughput" : "low_latency"); + } +#else // TODO(MaoZiming): improves pinning. // Offset LL-mode proxies onto a separate CPU range so a high-throughput // Buffer and an LL-mode Buffer running in the same process don't fight @@ -127,6 +176,7 @@ void Proxy::pin_thread_to_cpu_wrapper() { cpu, cfg_.thread_idx, cfg_.local_rank, cfg_.use_normal_mode ? "high_throughput" : "low_latency"); } +#endif } } @@ -188,6 +238,9 @@ void Proxy::init_common() { per_thread_rdma_init(ctx_, cfg_.gpu_buffer, cfg_.total_size, my_rank, cfg_.thread_idx, cfg_.local_rank); +#ifdef USE_CXI + pin_thread_to_cpu_wrapper(); +#else pin_thread_to_numa_wrapper(); if (!get_cq(ctx_)) { (void)create_per_thread_comp_channel(ctx_); @@ -219,6 +272,7 @@ void Proxy::init_common() { (unsigned long long)ctx_.atomic_buffer_mr->addr, (size_t)ctx_.atomic_buffer_mr->length, ctx_.atomic_buffer_mr->rkey); } +#endif if (ctxs_for_all_ranks_.empty()) { fprintf(stderr, @@ -232,6 +286,7 @@ void Proxy::init_common() { // NOTE: This must NOT alias cfg_.gpu_buffer (which is used for other // layouts). For IBV_WR_ATOMIC_FETCH_AND_ADD the NIC DMA-writes the old value // here. +#ifndef USE_CXI if (!ctx_.atomic_old_values_buf || !ctx_.atomic_old_values_mr) { size_t const atomic_buf_size = ProxyCtx::kMaxAtomicOps * sizeof(uint64_t); void* p = nullptr; @@ -252,8 +307,11 @@ void Proxy::init_common() { std::abort(); } } +#endif int num_ranks = ctxs_for_all_ranks_.size(); + int const normal_rank_stride = + normal_mode_rank_stride(num_ranks, cfg_.num_nodes); local_infos_.assign(num_ranks, RDMAConnectionInfo{}); remote_infos_.assign(num_ranks, RDMAConnectionInfo{}); @@ -270,7 +328,8 @@ void Proxy::init_common() { for (int p = 0; p < num_ranks; ++p) { if (p == my_rank) continue; if (peers_[p].ip == peers_[my_rank].ip) continue; - if (cfg_.use_normal_mode && std::abs(p - my_rank) % MAX_NUM_GPUS != 0) + if (cfg_.use_normal_mode && + std::abs(p - my_rank) % normal_rank_stride != 0) continue; ++num_active_peers; } @@ -296,6 +355,13 @@ void Proxy::init_common() { c.pd = ctx_.pd; c.mr = ctx_.mr; c.rkey = ctx_.rkey; +#ifdef USE_CXI + c.remote_addr = 0; + c.remote_len = cfg_.total_size; + c.numa_node = 0; + c.cxi_local_base = cfg_.gpu_buffer; + c.cxi_local_len = cfg_.total_size; +#endif #ifdef USE_DMABUF c.gpu_mr_chunks = ctx_.gpu_mr_chunks; #endif @@ -309,9 +375,11 @@ void Proxy::init_common() { c.atomic_old_values_mr = ctx_.atomic_old_values_mr; if (peer == my_rank) continue; - // Skip rdma connection for intra-node. - if (peers_[peer].ip == peers_[my_rank].ip) continue; - if (cfg_.use_normal_mode && std::abs(peer - my_rank) % MAX_NUM_GPUS != 0) + if (skip_network_peer_on_same_ip(peers_[peer].ip == peers_[my_rank].ip, + cfg_.use_normal_mode)) + continue; + if (cfg_.use_normal_mode && + std::abs(peer - my_rank) % normal_rank_stride != 0) continue; #ifdef EFA // Alias the shared SRD QPs from ctx_; dst_ah/dst_qpn (set later in @@ -325,6 +393,13 @@ void Proxy::init_common() { // Advertised local info is identical for all peers (shared QPs/MR/GID). local_infos_[peer] = template_local_info; #else +#ifdef USE_CXI + c.cxi_transport = std::make_unique(); + c.cxi_transport->init(cfg_.local_rank % 4); + c.cxi_transport->register_cuda_buffer(cfg_.gpu_buffer, cfg_.total_size); + c.cxi_local_base = cfg_.gpu_buffer; + c.cxi_local_len = cfg_.total_size; +#endif create_per_thread_qp(c, cfg_.gpu_buffer, cfg_.total_size, &local_infos_[peer], my_rank, cfg_.d2h_queues.size(), cfg_.use_normal_mode, atomic_buffer_ptr_); @@ -335,12 +410,13 @@ void Proxy::init_common() { usleep(50 * 1000); // Out-of-band exchange info per pair: start receiver thread first - std::thread receiver_thread([this, num_ranks, my_rank]() { + std::thread receiver_thread([this, num_ranks, my_rank, normal_rank_stride]() { for (int peer = 0; peer < num_ranks; ++peer) { - // Skip rdma connection for intra-node. - if (peer == my_rank || peers_[peer].ip == peers_[my_rank].ip || + if (peer == my_rank || + skip_network_peer_on_same_ip(peers_[peer].ip == peers_[my_rank].ip, + cfg_.use_normal_mode) || (cfg_.use_normal_mode && - std::abs(peer - my_rank) % MAX_NUM_GPUS != 0)) + std::abs(peer - my_rank) % normal_rank_stride != 0)) continue; int actual_peer; recv_connection_info_as_server(my_rank, &actual_peer, listen_fd_, @@ -350,8 +426,11 @@ void Proxy::init_common() { // Then send our info to all peers for (int peer = 0; peer < num_ranks; ++peer) { - if (peer == my_rank || peers_[peer].ip == peers_[my_rank].ip || - (cfg_.use_normal_mode && std::abs(peer - my_rank) % MAX_NUM_GPUS != 0)) + if (peer == my_rank || + skip_network_peer_on_same_ip(peers_[peer].ip == peers_[my_rank].ip, + cfg_.use_normal_mode) || + (cfg_.use_normal_mode && + std::abs(peer - my_rank) % normal_rank_stride != 0)) continue; char const* peer_ip = peers_[peer].ip.c_str(); int const peer_listen_port = peers_[peer].listen_ports[cfg_.thread_idx]; @@ -364,9 +443,13 @@ void Proxy::init_common() { // Verify remote info correctness for (int peer = 0; peer < num_ranks; ++peer) { - if (peer == my_rank || peers_[peer].ip == peers_[my_rank].ip || - (cfg_.use_normal_mode && std::abs(peer - my_rank) % MAX_NUM_GPUS != 0)) + if (peer == my_rank || + skip_network_peer_on_same_ip(peers_[peer].ip == peers_[my_rank].ip, + cfg_.use_normal_mode) || + (cfg_.use_normal_mode && + std::abs(peer - my_rank) % normal_rank_stride != 0)) continue; +#ifndef USE_CXI if (remote_infos_[peer].addr != peers_[peer].ptr) { fprintf(stderr, "Rank %d thread %d: Warning: remote addr mismatch for peer %d: " @@ -375,14 +458,17 @@ void Proxy::init_common() { peers_[peer].ptr); std::abort(); } +#endif } // Bring each per-peer QP to RTR/RTS for (int peer = 0; peer < num_ranks; ++peer) { if (peer == my_rank) continue; - // Skip rdma connection for intra-node. - if (peers_[peer].ip == peers_[my_rank].ip) continue; - if (cfg_.use_normal_mode && std::abs(peer - my_rank) % MAX_NUM_GPUS != 0) + if (skip_network_peer_on_same_ip(peers_[peer].ip == peers_[my_rank].ip, + cfg_.use_normal_mode)) + continue; + if (cfg_.use_normal_mode && + std::abs(peer - my_rank) % normal_rank_stride != 0) continue; auto& c = *ctxs_for_all_ranks_[peer]; @@ -393,6 +479,14 @@ void Proxy::init_common() { c.remote_addr = remote_infos_[peer].addr; c.remote_rkey = remote_infos_[peer].rkey; c.remote_len = remote_infos_[peer].len; +#ifdef USE_CXI + c.remote_addr = 0; + c.remote_len = remote_infos_[peer].cxi_mr_len; + c.cxi_remote_key = remote_infos_[peer].cxi_mr_key; + c.cxi_remote_len = remote_infos_[peer].cxi_mr_len; + c.cxi_remote_host_key = remote_infos_[peer].cxi_host_mr_key; + c.cxi_remote_host_len = remote_infos_[peer].cxi_host_mr_len; +#endif #ifdef USE_DMABUF // Populate remote MR chunks from exchanged connection info. @@ -421,6 +515,7 @@ void Proxy::init_common() { } #endif +#ifndef USE_CXI // Set remote atomic buffer info from exchanged connection info c.remote_atomic_buffer_addr = remote_infos_[peer].atomic_buffer_addr; c.remote_atomic_buffer_len = remote_infos_[peer].atomic_buffer_len; @@ -439,6 +534,7 @@ void Proxy::init_common() { peer, (unsigned long long)c.remote_atomic_buffer_addr, (size_t)c.remote_atomic_buffer_len, c.remote_atomic_buffer_rkey); } +#endif } usleep(50 * 1000); if (cfg_.use_normal_mode) { @@ -500,7 +596,7 @@ void Proxy::init_sender() { init_common(); assert(cfg_.rank == 0); auto& ctx_ptr = ctxs_for_all_ranks_[1]; -#ifndef EFA +#if !defined(EFA) && !defined(USE_CXI) local_post_ack_buf(*ctx_ptr, kSenderAckQueueDepth); #else // EFA: posted once on the shared recv_ack_qp in init_common. @@ -512,12 +608,14 @@ void Proxy::init_remote() { init_common(); assert(cfg_.rank == 1); auto& ctx_ptr = ctxs_for_all_ranks_[0]; -#ifndef EFA +#if !defined(EFA) && !defined(USE_CXI) local_post_ack_buf(*ctx_ptr, kSenderAckQueueDepth); #endif +#ifndef USE_CXI remote_reg_ack_buf(ctx_ptr->pd, ring.ack_buf, ring.ack_mr); ring.ack_qp = ctx_ptr->ack_qp; -#ifndef EFA +#endif +#if !defined(EFA) && !defined(USE_CXI) post_receive_buffer_for_imm(*ctx_ptr); #endif } @@ -554,20 +652,27 @@ void Proxy::run_remote() { void Proxy::run_dual() { init_common(); + int const normal_rank_stride = normal_mode_rank_stride( + static_cast(ctxs_for_all_ranks_.size()), cfg_.num_nodes); for (int peer = 0; peer < (int)ctxs_for_all_ranks_.size(); ++peer) { if (peer == cfg_.rank) continue; - if (peers_[peer].ip == peers_[cfg_.rank].ip) continue; - if (cfg_.use_normal_mode && std::abs(peer - cfg_.rank) % MAX_NUM_GPUS != 0) + if (skip_network_peer_on_same_ip(peers_[peer].ip == peers_[cfg_.rank].ip, + cfg_.use_normal_mode)) + continue; + if (cfg_.use_normal_mode && + std::abs(peer - cfg_.rank) % normal_rank_stride != 0) continue; auto& ctx_ptr = ctxs_for_all_ranks_[peer]; if (!ctx_ptr) continue; -#ifndef EFA +#if !defined(EFA) && !defined(USE_CXI) // EFA: posted once on the shared recv_ack_qp in init_common. local_post_ack_buf(*ctx_ptr, kSenderAckQueueDepth); #endif +#ifndef USE_CXI remote_reg_ack_buf(ctx_ptr->pd, ring.ack_buf, ring.ack_mr); ring.ack_qp = ctx_ptr->ack_qp; -#ifndef EFA +#endif +#if !defined(EFA) && !defined(USE_CXI) post_receive_buffer_for_imm(*ctx_ptr); #endif } @@ -582,6 +687,16 @@ void Proxy::run_dual() { atomic_buffer_ptr_, cfg_.num_ranks, cfg_.num_experts, pending_atomic_updates, cfg_.rank, cfg_.num_nodes, adaptive_sleeper_, cfg_.use_normal_mode); +#ifdef USE_CXI + if (cxi_fast_mode()) { + for (auto& c : ctxs_for_all_ranks_) { + if (c && c->cxi_transport && + (c->cxi_transport->outstanding() > 0 || + !c->cxi_pending_atomics.empty())) + cxi_retire_ctx(*c, acked_wrs_, 64); + } + } +#endif notify_gpu_completion(my_tail); post_gpu_command(my_tail, seen); #ifdef USE_RECEIVER_BARRIER @@ -607,6 +722,23 @@ void Proxy::run_dual() { barrier_check(); } } + +#ifdef USE_CXI + // Best-effort drain so transports are not torn down with writes in + // flight; bounded so shutdown cannot hang on a dead peer. + if (cxi_fast_mode()) { + auto const deadline = + std::chrono::steady_clock::now() + std::chrono::milliseconds(200); + for (auto& c : ctxs_for_all_ranks_) { + if (!c || !c->cxi_transport) continue; + while ((c->cxi_transport->outstanding() > 0 || + !c->cxi_pending_atomics.empty()) && + std::chrono::steady_clock::now() < deadline) { + cxi_retire_ctx(*c, acked_wrs_, 64); + } + } + } +#endif } void Proxy::notify_gpu_completion(uint64_t& my_tail) { @@ -673,6 +805,7 @@ void Proxy::post_gpu_command(uint64_t& my_tail, size_t& seen) { // Multi-ring buffer processing: collect commands from all ring buffers // Process each ring buffer (similar to test_multi_ring_throughput.cu) for (size_t rb_idx = 0; rb_idx < cfg_.d2h_queues.size(); rb_idx++) { + if (!ctx_.progress_run.load(std::memory_order_acquire)) return; d2hq::HostD2HHandle* h = &cfg_.d2h_queues[rb_idx]; #ifdef USE_MSCCLPP_FIFO_BACKEND assert(h && "h is empty!\n"); @@ -760,6 +893,7 @@ void Proxy::post_gpu_command(uint64_t& my_tail, size_t& seen) { // Collect batch of commands from this ring buffer for (size_t i = ring_seen; i < cur_head; ++i) { + if (!ctx_.progress_run.load(std::memory_order_acquire)) return; CmdType cmd = h->volatile_load_cmd_type(i); // NOTE(MaoZiming): Non-blocking. prevent local and remote both while // loop. @@ -779,7 +913,9 @@ void Proxy::post_gpu_command(uint64_t& my_tail, size_t& seen) { cudaDeviceSynchronize(); std::abort(); } - if (peers_[cmd_entry.dst_rank].ip == peers_[cfg_.rank].ip) { + if (skip_network_peer_on_same_ip( + peers_[cmd_entry.dst_rank].ip == peers_[cfg_.rank].ip, + cfg_.use_normal_mode)) { fprintf( stderr, "[ERROR] Intra-node command!, cmd.dst_rank: %d, cfg_.rank: %d, " @@ -806,6 +942,7 @@ void Proxy::post_gpu_command(uint64_t& my_tail, size_t& seen) { // Process all collected commands in batch if (!wrs_to_post.empty()) { + if (!ctx_.progress_run.load(std::memory_order_acquire)) return; #ifdef MEASURE_PER_OP_LATENCY auto start = std::chrono::high_resolution_clock::now(); #endif @@ -1025,7 +1162,7 @@ void Proxy::post_gpu_commands_mixed( if (!rdma_wrs.empty()) { post_rdma_async_batched(ctx_, cfg_.gpu_buffer, rdma_wrs.size(), rdma_wrs, rdma_cmds, ctxs_for_all_ranks_, cfg_.rank, - cfg_.thread_idx, cfg_.use_normal_mode); + cfg_.thread_idx, cfg_.use_normal_mode, acked_wrs_); rdma_wrs.clear(); rdma_cmds.clear(); } @@ -1060,6 +1197,23 @@ void Proxy::post_gpu_commands_mixed( } void Proxy::quiet_cq() { +#ifdef USE_CXI + if (cxi_fast_mode()) { + // Drain all outstanding async writes AND deferred control atomics so + // quiet means "everything sent", matching the verbs path's semantics. + for (auto& c : ctxs_for_all_ranks_) { + if (!c || !c->cxi_transport) continue; + while (c->cxi_transport->outstanding() > 0 || + !c->cxi_pending_atomics.empty()) { + if (cxi_retire_ctx(*c, acked_wrs_, 64) == 0) { + if (!ctx_.progress_run.load(std::memory_order_acquire)) return; + cpu_relax(); + } + } + } + } + return; +#endif auto outstanding_batches = [&]() -> size_t { return 0; }; constexpr int kConsecutiveEmptyToExit = 3; int empty_iters = 0; @@ -1305,6 +1459,38 @@ void Proxy::destroy(bool free_gpu_buffer) { void Proxy::post_barrier_msg(int dst_rank, bool ack, uint64_t seq) { ProxyCtx* ctx = ctxs_for_all_ranks_[dst_rank].get(); +#ifdef USE_CXI + (void)seq; + if (!ctx || !ctx->cxi_transport || ctx->cxi_peer_addr == FI_ADDR_UNSPEC || + ctx->cxi_remote_host_key == 0) { + fprintf(stderr, "barrier_msg: bad CXI ctx for dst=%d\n", dst_rank); + std::abort(); + } + int const stride = normal_mode_rank_stride(cfg_.num_ranks, cfg_.num_nodes); + int const dst_node_idx = stride > 0 ? dst_rank / stride : 0; + int const slot = ack ? cfg_.num_nodes + dst_node_idx : cfg_.node_idx; + if (cxi_fast_mode()) { + // Barriers imply all prior traffic to this peer is on the wire: drain + // outstanding writes and deferred atomics first (barriers are rare). + while (ctx->cxi_transport->outstanding() > 0 || + !ctx->cxi_pending_atomics.empty()) { + if (cxi_retire_ctx(*ctx, acked_wrs_, 64) == 0) { + if (!ctx_.progress_run.load(std::memory_order_acquire)) return; + cpu_relax(); + } + } + while (ctx->cxi_transport->try_inject_atomic_add64( + ctx->cxi_peer_addr, 1, cxi_barrier_slot_offset(slot), + ctx->cxi_remote_host_key) == -FI_EAGAIN) { + if (cxi_retire_ctx(*ctx, acked_wrs_, 64) == 0) cpu_relax(); + } + } else { + ctx->cxi_transport->inject_atomic_add64(ctx->cxi_peer_addr, 1, + cxi_barrier_slot_offset(slot), + ctx->cxi_remote_host_key); + } + return; +#endif if (!ctx || !ctx->qp || !ctx->mr) { fprintf(stderr, "barrier_msg: bad ctx for dst=%d\n", dst_rank); std::abort(); @@ -1362,7 +1548,17 @@ void Proxy::send_barrier(uint64_t wr) { #endif assert(ctx_.barrier_wr == -1 && "barrier_wr should be 0"); ctx_.barrier_wr = wr; +#ifdef USE_CXI + // The CXI barrier compares seq against per-node slot counters that are + // NIC atomic adds and never reset, so seq must stay monotonic for the + // process lifetime: wrapping at kSeqMask (2^21) would make every + // `slots[node] >= seq` check pass spuriously and release barriers + // without synchronizing. The 21-bit mask only exists for the verbs + // immediate-data encoding, which the CXI path does not use. + ctx_.barrier_seq = ctx_.barrier_seq + 1; +#else ctx_.barrier_seq = (ctx_.barrier_seq + 1) & BarrierImm::kSeqMask; +#endif if (cfg_.rank == ctx_.node_leader_rank) { if (ctx_.barrier_arrived.size() != static_cast(cfg_.num_nodes)) { @@ -1402,12 +1598,13 @@ void Proxy::barrier_check() { ++ctx_.barrier_arrival_count; } } else { - int rank = cfg_.rank - cfg_.node_idx * MAX_NUM_GPUS; - if (rank < 0 || rank >= MAX_NUM_GPUS) { + int const local_stride = std::max(1, ctx_.num_local_ranks); + int rank = cfg_.rank - cfg_.node_idx * local_stride; + if (rank < 0 || rank >= local_stride) { printf("rank: %d, node_idx: %d invalid for barrier\n", cfg_.rank, cfg_.node_idx); } - assert(rank >= 0 && rank < MAX_NUM_GPUS); + assert(rank >= 0 && rank < local_stride); post_barrier_msg(/*dst=*/rank, /*ack=*/false, seq); } @@ -1417,7 +1614,8 @@ void Proxy::barrier_check() { if (ctx_.barrier_arrival_count == cfg_.num_nodes) { std::unordered_map leader_for_ip; for (int r = 0; r < (int)peers_.size(); ++r) { - if (r >= MAX_NUM_GPUS && (r - cfg_.rank) % MAX_NUM_GPUS == 0) { + int const local_stride = std::max(1, ctx_.num_local_ranks); + if (r >= local_stride && (r - cfg_.rank) % local_stride == 0) { leader_for_ip[peers_[r].ip] = r; } } @@ -1473,6 +1671,61 @@ void Proxy::barrier_check() { } if (all_local_arrived) { static thread_local uint64_t last_sent_seq = 0; +#ifdef USE_CXI + auto* slots = cxi_barrier_slots(atomic_buffer_ptr_); + if (last_sent_seq != seq) { + last_sent_seq = seq; + if (cfg_.rank != 0) { + post_barrier_msg(/*dst=*/0, /*ack=*/false, seq); + } + } + if (cfg_.rank == 0) { + bool all_nodes_arrived = true; + for (int node = 1; node < cfg_.num_nodes; ++node) { + if (slots[node].load(std::memory_order_acquire) < + static_cast(seq)) { + all_nodes_arrived = false; + break; + } + } + if (all_nodes_arrived) { + std::unordered_map leader_for_ip; + for (int r = 0; r < (int)peers_.size(); ++r) { + auto it = leader_for_ip.find(peers_[r].ip); + if (it == leader_for_ip.end() || r < it->second) { + assert(ctx_.num_local_ranks <= 0 || + r % ctx_.num_local_ranks == 0); + leader_for_ip[peers_[r].ip] = r; + } + } + for (auto const& kv : leader_for_ip) { + std::string const& ip = kv.first; + int leader_r = kv.second; + if (ip == peers_[0].ip) continue; + post_barrier_msg(leader_r, true, seq); + } + for (int lr = 0; lr < ctx_.num_local_ranks; ++lr) { + lb->release_seq[lr].store(seq, std::memory_order_release); + } + acked_wrs_.insert(ctx_.barrier_wr); +#ifndef USE_MSCCLPP_FIFO_BACKEND + ctx_.barrier_inflight = false; + ctx_.barrier_wr = -1; +#endif + return; + } + } else if (slots[cfg_.num_nodes + cfg_.node_idx].load( + std::memory_order_acquire) >= static_cast(seq)) { + for (int lr = 0; lr < ctx_.num_local_ranks; ++lr) { + lb->release_seq[lr].store(seq, std::memory_order_release); + } + acked_wrs_.insert(ctx_.barrier_wr); +#ifndef USE_MSCCLPP_FIFO_BACKEND + ctx_.barrier_inflight = false; + ctx_.barrier_wr = -1; +#endif + } +#else if (last_sent_seq != seq) { last_sent_seq = seq; if (cfg_.rank == 0) { @@ -1497,7 +1750,7 @@ void Proxy::barrier_check() { for (int r = 0; r < (int)peers_.size(); ++r) { auto it = leader_for_ip.find(peers_[r].ip); if (it == leader_for_ip.end() || r < it->second) { - assert(r % MAX_NUM_GPUS == 0); + assert(ctx_.num_local_ranks <= 0 || r % ctx_.num_local_ranks == 0); leader_for_ip[peers_[r].ip] = r; } } @@ -1539,6 +1792,7 @@ void Proxy::barrier_check() { ctx_.barrier_wr = -1; #endif } +#endif } return; } else { @@ -1555,4 +1809,4 @@ void Proxy::barrier_check() { #endif } } -#endif \ No newline at end of file +#endif diff --git a/ep/src/rdma.cpp b/ep/src/rdma.cpp index e01374673..9e9ddeeef 100644 --- a/ep/src/rdma.cpp +++ b/ep/src/rdma.cpp @@ -28,6 +28,7 @@ #include #include #include +#include #include #include #include @@ -391,6 +392,22 @@ bool is_cuda_host_pointer(void* ptr) { void per_thread_rdma_init(ProxyCtx& S, void* gpu_buf, size_t bytes, int rank, int thread_idx, int local_rank) { +#ifdef USE_CXI + if (S.cxi_transport) return; + (void)rank; + (void)thread_idx; + cudaSetDevice(local_rank); + S.cxi_transport = std::make_unique(); + S.cxi_transport->init(local_rank % 4); + S.cxi_transport->register_cuda_buffer(gpu_buf, bytes); + S.cxi_local_base = gpu_buf; + S.cxi_local_len = bytes; + S.remote_addr = 0; + S.remote_len = bytes; + S.numa_node = 0; + return; +#endif + if (S.context) return; // already initialized int num_devices = 0; @@ -981,6 +998,36 @@ void create_per_thread_qp(ProxyCtx& S, void* gpu_buffer, size_t size, RDMAConnectionInfo* local_info, int rank, size_t num_rings, bool use_normal_mode, void* atomic_buffer_ptr) { +#ifdef USE_CXI + (void)gpu_buffer; + (void)rank; + (void)num_rings; + (void)use_normal_mode; + if (!S.cxi_transport) { + fprintf(stderr, "CXI transport is not initialized\n"); + std::abort(); + } + if (atomic_buffer_ptr) { + S.cxi_transport->register_host_buffer(atomic_buffer_ptr, kAtomicBufferSize); + } + auto info = S.cxi_transport->local_info(); + if (info.ep_name.size() > sizeof(local_info->cxi_ep_name)) { + fprintf(stderr, "CXI endpoint name too large: %zu > %zu\n", + info.ep_name.size(), sizeof(local_info->cxi_ep_name)); + std::abort(); + } + local_info->addr = 0; + local_info->len = size; + local_info->cxi_mr_key = info.mr_key; + local_info->cxi_mr_len = info.size; + local_info->cxi_host_mr_key = info.host_mr_key; + local_info->cxi_host_mr_len = info.host_size; + local_info->cxi_ep_name_len = static_cast(info.ep_name.size()); + std::memcpy(local_info->cxi_ep_name, info.ep_name.data(), + info.ep_name.size()); + return; +#endif + if (S.qp) return; // Already initialized for this thread if (S.ack_qp) return; if (S.recv_ack_qp) return; @@ -1100,7 +1147,7 @@ void create_per_thread_qp(ProxyCtx& S, void* gpu_buffer, size_t size, } void modify_qp_to_init(ProxyCtx& S) { -#ifdef EFA +#if defined(EFA) || defined(USE_CXI) return; #endif struct ibv_qp_attr attr; @@ -1168,6 +1215,27 @@ struct ibv_ah* create_ah(ProxyCtx& S, uint8_t* remote_gid) { void modify_qp_to_rtr(ProxyCtx& S, RDMAConnectionInfo* remote, bool use_normal_mode) { +#ifdef USE_CXI + (void)use_normal_mode; + if (!S.cxi_transport) { + fprintf(stderr, "CXI transport is not initialized for peer\n"); + std::abort(); + } + uccl::cxi::EndpointInfo peer_info; + peer_info.mr_key = remote->cxi_mr_key; + peer_info.size = remote->cxi_mr_len; + peer_info.host_mr_key = remote->cxi_host_mr_key; + peer_info.host_size = remote->cxi_host_mr_len; + peer_info.ep_name.assign(remote->cxi_ep_name, + remote->cxi_ep_name + remote->cxi_ep_name_len); + S.cxi_peer_addr = S.cxi_transport->insert_peer(peer_info); + S.cxi_remote_key = remote->cxi_mr_key; + S.cxi_remote_len = remote->cxi_mr_len; + S.cxi_remote_host_key = remote->cxi_host_mr_key; + S.cxi_remote_host_len = remote->cxi_host_mr_len; + return; +#endif + #ifdef EFA S.dst_qpn = remote->qp_num; S.dst_ack_qpn = remote->recv_ack_qp_num; @@ -1303,6 +1371,12 @@ void modify_qp_to_rtr(ProxyCtx& S, RDMAConnectionInfo* remote, } void modify_qp_to_rts(ProxyCtx& S, RDMAConnectionInfo* local_info) { +#ifdef USE_CXI + (void)S; + (void)local_info; + return; +#endif + #ifdef EFA return; #endif @@ -1384,6 +1458,291 @@ void post_receive_buffer_for_imm(ProxyCtx& S) { } // Normal mode implementation +#ifdef USE_CXI +struct CxiPlannedWrite { + uccl::cxi::Transport* transport = nullptr; + ProxyCtx* ctx = nullptr; + fi_addr_t peer = FI_ADDR_UNSPEC; + void* local = nullptr; + uint64_t local_offset = 0; + uint64_t remote_offset = 0; + uint64_t remote_key = 0; + size_t bytes = 0; + std::vector wr_ids; +}; + +// Fast (default): writes are posted asynchronously and ring slots retire on +// CQ completion; tail/control atomics are injected immediately behind their +// writes, relying on same-TX-context submission ordering (validated +// empirically on Slingshot-11, see uccl-project/uccl#956). Set +// UCCL_CXI_SYNC_WRITES=1 to restore the conservative +// post/wait-all/then-atomics behavior. +bool cxi_fast_mode() { + static bool const fast = std::getenv("UCCL_CXI_SYNC_WRITES") == nullptr; + return fast; +} + +static size_t cxi_max_outstanding() { + static size_t const cap = [] { + char const* v = std::getenv("UCCL_CXI_MAX_OUTSTANDING"); + return v ? std::max(1ul, strtoul(v, nullptr, 10)) : 512ul; + }(); + return cap; +} + +// Poll one peer's transport, retiring completed writes into acked_wrs and +// returning contexts to the pool, then inject any control atomics whose +// gating writes have all completed. Returns completions consumed. +size_t cxi_retire_ctx(ProxyCtx& ctx, std::unordered_set& acked_wrs, + size_t budget) { + if (!ctx.cxi_transport) return 0; + auto* transport = ctx.cxi_transport.get(); + size_t consumed = 0; + uccl::cxi::WriteContext* done[16]; + while (consumed < budget) { + size_t const n = + transport->poll(done, std::min(16, budget - consumed)); + if (n == 0) break; + for (size_t i = 0; i < n; ++i) { + for (uint64_t wr : done[i]->wr_ids) acked_wrs.insert(wr); + transport->release_context(done[i]); + } + ctx.cxi_writes_completed += n; + consumed += n; + } + + // Flush ripe pending atomics in order. + while (!ctx.cxi_pending_atomics.empty() && + ctx.cxi_pending_atomics.front().threshold <= + ctx.cxi_writes_completed) { + auto const& pa = ctx.cxi_pending_atomics.front(); + while (transport->try_inject_atomic_add64(ctx.cxi_peer_addr, pa.value, + pa.remote_offset, + ctx.cxi_remote_host_key) == + -FI_EAGAIN) { + // TX full: consume more completions to free credits. + size_t const n = transport->poll(done, 16); + for (size_t i = 0; i < n; ++i) { + for (uint64_t wr : done[i]->wr_ids) acked_wrs.insert(wr); + transport->release_context(done[i]); + } + ctx.cxi_writes_completed += n; + if (n == 0) cpu_relax(); + } + ctx.cxi_pending_atomics.pop_front(); + } + return consumed; +} + +// True when nothing to this peer is still in flight or deferred. +static bool cxi_ctx_idle(ProxyCtx& ctx) { + return (!ctx.cxi_transport || ctx.cxi_transport->outstanding() == 0) && + ctx.cxi_pending_atomics.empty(); +} + +static void maybe_print_cxi_write_stats( + int my_rank, std::vector const& wrs_to_post, + std::vector const& cmds_to_post, + std::vector const& writes) { + static bool const enabled = std::getenv("UCCL_CXI_STATS") != nullptr; + if (!enabled || cmds_to_post.empty()) return; + + size_t total_bytes = 0; + size_t min_bytes = SIZE_MAX; + size_t max_bytes = 0; + size_t dispatch_cmds = 0; + size_t combine_cmds = 0; + std::array buckets{}; + for (auto const& cmd : cmds_to_post) { + if (get_is_combine(cmd.cmd_type)) { + ++combine_cmds; + } else { + ++dispatch_cmds; + } + } + for (auto const& write : writes) { + total_bytes += write.bytes; + min_bytes = std::min(min_bytes, write.bytes); + max_bytes = std::max(max_bytes, write.bytes); + size_t bucket = 0; + size_t n = std::max(1, write.bytes); + while (n > 4096 && bucket + 1 < buckets.size()) { + n >>= 1; + ++bucket; + } + ++buckets[bucket]; + } + if (writes.empty()) min_bytes = 0; + + static std::atomic seq{0}; + uint64_t const id = seq.fetch_add(1, std::memory_order_relaxed); + char const* phase = std::getenv("UCCL_CXI_PHASE"); + if (!phase) phase = "unset"; + fprintf(stderr, + "[CXI_STATS] id=%llu rank=%d phase=%s kind=%s cmds=%zu dispatch_cmds=%zu " + "combine_cmds=%zu writes=%zu bytes=%zu min=%zu max=%zu " + "avg=%.1f coalesced=%.2f buckets_le4k=%zu " + "le8k=%zu le16k=%zu le32k=%zu le64k=%zu le128k=%zu le256k=%zu " + "gt256k=%zu\n", + (unsigned long long)id, my_rank, phase, + combine_cmds != 0 && dispatch_cmds == 0 ? "combine" : "dispatch", + wrs_to_post.size(), dispatch_cmds, combine_cmds, writes.size(), + total_bytes, min_bytes, max_bytes, + writes.empty() ? 0.0 + : static_cast(total_bytes) / + static_cast(writes.size()), + writes.empty() + ? 0.0 + : static_cast(cmds_to_post.size()) / + static_cast(writes.size()), + buckets[0], buckets[1], buckets[2], + buckets[3], buckets[4], buckets[5], buckets[6], buckets[7]); +} + +static void post_rdma_async_batched_cxi( + std::vector const& wrs_to_post, + std::vector const& cmds_to_post, + std::vector>& ctxs, int my_rank, + bool use_normal_mode, std::atomic const* progress_run, + std::unordered_set& acked_wrs) { + std::vector writes; + writes.reserve(cmds_to_post.size()); + + auto append_write = [&writes](uccl::cxi::Transport* transport, ProxyCtx* ctx, + fi_addr_t peer, void* local, + uint64_t local_offset, + uint64_t remote_offset, uint64_t remote_key, + size_t bytes, uint64_t wr_id) { + if (!writes.empty()) { + auto& prev = writes.back(); + uintptr_t const prev_local_end = + reinterpret_cast(prev.local) + prev.bytes; + if (prev.transport == transport && prev.ctx == ctx && prev.peer == peer && + prev.remote_key == remote_key && + prev.local_offset + prev.bytes == local_offset && + prev.remote_offset + prev.bytes == remote_offset && + prev_local_end == reinterpret_cast(local)) { + prev.bytes += bytes; + prev.wr_ids.push_back(wr_id); + return; + } + } + writes.push_back(CxiPlannedWrite{transport, ctx, peer, local, local_offset, + remote_offset, remote_key, bytes, + {wr_id}}); + }; + + for (size_t i = 0; i < wrs_to_post.size(); ++i) { + auto const& cmd = cmds_to_post[i]; + if (cmd.dst_rank == static_cast(my_rank)) { + fprintf(stderr, "Posting CXI write to itself\n"); + std::abort(); + } + ProxyCtx* ctx = ctxs[cmd.dst_rank].get(); + if (!ctx || !ctx->cxi_transport || + ctx->cxi_peer_addr == FI_ADDR_UNSPEC || ctx->cxi_remote_key == 0) { + fprintf(stderr, + "Destination CXI ctx missing fields: src=%d dst=%u " + "ctx=%p transport=%p peer_addr=%llu " + "remote_key=%llu remote_len=%llu\n", + my_rank, cmd.dst_rank, static_cast(ctx), + ctx ? static_cast(ctx->cxi_transport.get()) : nullptr, + ctx ? static_cast(ctx->cxi_peer_addr) : 0ULL, + ctx ? static_cast(ctx->cxi_remote_key) : 0ULL, + ctx ? static_cast(ctx->cxi_remote_len) : 0ULL); + std::abort(); + } + uint64_t const remote_offset = + decode_write_offset(cmd.req_rptr, !use_normal_mode); + uint64_t const local_offset = + decode_write_offset(cmd.req_lptr, !use_normal_mode); + if (remote_offset + cmd.bytes > ctx->cxi_remote_len) { + fprintf(stderr, + "[CXI] Remote write OOB: offset=0x%llx len=%u size=%zu\n", + (unsigned long long)remote_offset, cmd.bytes, + (size_t)ctx->cxi_remote_len); + std::abort(); + } + if (!ctx->cxi_local_base) { + fprintf(stderr, "CXI local base is missing for dst=%u\n", cmd.dst_rank); + std::abort(); + } + if (local_offset + cmd.bytes > ctx->cxi_local_len) { + fprintf(stderr, + "[CXI] Local write OOB: rank=%d dst=%u offset=0x%llx len=%u " + "size=%zu normal=%d\n", + my_rank, cmd.dst_rank, (unsigned long long)local_offset, + cmd.bytes, (size_t)ctx->cxi_local_len, + static_cast(use_normal_mode)); + std::abort(); + } + void* local = reinterpret_cast( + reinterpret_cast(ctx->cxi_local_base) + local_offset); + (void)use_normal_mode; + append_write(ctx->cxi_transport.get(), ctx, ctx->cxi_peer_addr, local, + local_offset, remote_offset, ctx->cxi_remote_key, cmd.bytes, + wrs_to_post[i]); + } + + maybe_print_cxi_write_stats(my_rank, wrs_to_post, cmds_to_post, writes); + + if (cxi_fast_mode()) { + // Async path: post and return. Ring slots retire via cxi_retire_ctx when + // completions surface; ordering vs the tail atomics that follow comes + // from same-TX-context submission order. + size_t const cap = cxi_max_outstanding(); + for (auto& write : writes) { + while (write.transport->outstanding() >= cap) { + if (cxi_retire_ctx(*write.ctx, acked_wrs, 64) == 0) { + if (progress_run && + !progress_run->load(std::memory_order_acquire)) + return; + cpu_relax(); + } + } + uccl::cxi::WriteContext* c = write.transport->acquire_context(); + c->wr_ids = std::move(write.wr_ids); + while (write.transport->try_write(write.peer, write.local, write.bytes, + write.remote_offset, write.remote_key, + c) == -FI_EAGAIN) { + if (cxi_retire_ctx(*write.ctx, acked_wrs, 64) == 0) { + if (progress_run && + !progress_run->load(std::memory_order_acquire)) { + write.transport->release_context(c); + return; + } + cpu_relax(); + } + } + ++write.ctx->cxi_writes_posted; + } + return; + } + + // Conservative path: post everything, wait for all completions, retire in + // bulk. Guarantees write completion (NIC-acked) before the caller posts + // the corresponding atomics. + std::vector write_contexts(writes.size()); + std::unordered_map> + pending_by_transport; + pending_by_transport.reserve(writes.size()); + + for (size_t i = 0; i < writes.size(); ++i) { + auto const& write = writes[i]; + write.transport->write(write.peer, write.local, write.bytes, + write.remote_offset, write.remote_key, + &write_contexts[i]); + pending_by_transport[write.transport].push_back(&write_contexts[i]); + } + + for (auto& [transport, pending] : pending_by_transport) { + if (!transport->wait_all(pending, progress_run)) return; + } + for (uint64_t wr : wrs_to_post) acked_wrs.insert(wr); +} +#endif + static void post_rdma_async_batched_normal_mode( ProxyCtx& S, void* buf, size_t num_wrs, std::vector const& wrs_to_post, @@ -2078,7 +2437,18 @@ void post_rdma_async_batched(ProxyCtx& S, void* buf, size_t num_wrs, std::vector const& cmds_to_post, std::vector>& ctxs, int my_rank, int thread_idx, - bool use_normal_mode) { + bool use_normal_mode, + std::unordered_set& acked_wrs) { +#ifdef USE_CXI + (void)S; + (void)buf; + (void)num_wrs; + (void)thread_idx; + post_rdma_async_batched_cxi(wrs_to_post, cmds_to_post, ctxs, my_rank, + use_normal_mode, &S.progress_run, acked_wrs); + return; +#endif + (void)acked_wrs; if (use_normal_mode) { post_rdma_async_batched_normal_mode( S, buf, num_wrs, wrs_to_post, cmds_to_post, ctxs, my_rank, thread_idx); @@ -3461,6 +3831,68 @@ void post_atomic_operations(ProxyCtx& S, int my_rank, int thread_idx, std::unordered_set& acked_wrs, bool use_normal_mode) { +#ifdef USE_CXI + (void)thread_idx; + (void)use_normal_mode; + if (wrs_to_post.size() != cmds_to_post.size()) { + fprintf(stderr, "CXI atomic size mismatch: wrs=%zu cmds=%zu\n", + wrs_to_post.size(), cmds_to_post.size()); + std::abort(); + } + for (size_t i = 0; i < cmds_to_post.size(); ++i) { + if (!S.progress_run.load(std::memory_order_acquire)) return; + auto const& cmd = cmds_to_post[i]; + if (cmd.dst_rank == static_cast(my_rank)) { + fprintf(stderr, "Posting CXI atomic to itself\n"); + std::abort(); + } + ProxyCtx* ctx = ctxs[cmd.dst_rank].get(); + if (!ctx || !ctx->cxi_transport || + ctx->cxi_peer_addr == FI_ADDR_UNSPEC || ctx->cxi_remote_host_key == 0) { + fprintf(stderr, "Destination CXI host MR missing for dst=%u\n", + cmd.dst_rank); + std::abort(); + } + if ((cmd.req_rptr & 0x7) != 0 || + static_cast(cmd.req_rptr) + sizeof(int64_t) > + ctx->cxi_remote_host_len) { + fprintf(stderr, + "[CXI] Remote atomic OOB or unaligned: offset=0x%x len=%zu\n", + cmd.req_rptr, (size_t)ctx->cxi_remote_host_len); + std::abort(); + } + int v = static_cast(cmd.value); + if (get_is_combine(cmd.cmd_type)) v = 1; + if (v == kLargeAtomicValue) v = kMaxSendAtomicValue; + int64_t const value64 = static_cast(static_cast(v)); + if (cxi_fast_mode()) { + // Control atomics must not overtake the data writes posted before + // them. If writes to this peer are still in flight (or earlier + // atomics are still queued), defer; cxi_retire_ctx injects them as + // their gating writes complete. + if (!ctx->cxi_pending_atomics.empty() || + ctx->cxi_writes_completed < ctx->cxi_writes_posted) { + ctx->cxi_pending_atomics.push_back(ProxyCtx::CxiPendingAtomic{ + value64, cmd.req_rptr, ctx->cxi_writes_posted}); + } else { + while (ctx->cxi_transport->try_inject_atomic_add64( + ctx->cxi_peer_addr, value64, cmd.req_rptr, + ctx->cxi_remote_host_key) == -FI_EAGAIN) { + if (cxi_retire_ctx(*ctx, acked_wrs, 64) == 0) { + if (!S.progress_run.load(std::memory_order_acquire)) return; + cpu_relax(); + } + } + } + } else { + ctx->cxi_transport->inject_atomic_add64(ctx->cxi_peer_addr, value64, + cmd.req_rptr, + ctx->cxi_remote_host_key); + } + acked_wrs.insert(wrs_to_post[i]); + } + return; +#endif if (use_normal_mode) { #ifndef EFA post_atomic_operations_native_rdma(S, wrs_to_post, cmds_to_post, ctxs, diff --git a/ep/src/uccl_ep.cc b/ep/src/uccl_ep.cc index d539163df..9351f8449 100644 --- a/ep/src/uccl_ep.cc +++ b/ep/src/uccl_ep.cc @@ -24,6 +24,7 @@ #include #include #include +#include #include #include #include @@ -542,7 +543,7 @@ class Buffer { if (num_tokens_per_rdma_rank_ptr != 0) { CUDA_CHECK(cudaMemsetAsync( reinterpret_cast(num_tokens_per_rdma_rank_ptr), 0, - num_ranks * sizeof(int), comm_stream)); + num_rdma_ranks * sizeof(int), comm_stream)); } CUDA_CHECK( cudaMemsetAsync(reinterpret_cast(num_tokens_per_expert_ptr), 0, diff --git a/ep/src/uccl_proxy.cpp b/ep/src/uccl_proxy.cpp index 92d434b8a..6b3d7da4d 100644 --- a/ep/src/uccl_proxy.cpp +++ b/ep/src/uccl_proxy.cpp @@ -5,9 +5,11 @@ #include "ring_buffer.cuh" #include #include +#include #include #include #include +#include #include UcclProxy::UcclProxy(int thread_idx, uintptr_t gpu_buffer_addr, @@ -57,32 +59,72 @@ UcclProxy::UcclProxy(int thread_idx, uintptr_t gpu_buffer_addr, node_idx_ = node_idx; if (thread_idx == 0) { -#ifdef USE_GRACE_HOPPER - cudaMallocManaged(&atomic_buffer_ptr_, kAtomicBufferSize); +#ifdef USE_CXI + auto err = cudaHostAlloc(&atomic_buffer_ptr_, kAtomicBufferSize, + cudaHostAllocMapped); + if (err != cudaSuccess) { + throw std::runtime_error(std::string("cudaHostAlloc atomic buffer failed: ") + + cudaGetErrorString(err)); + } + atomic_buffer_is_host_allocated_ = true; +#elif defined(USE_GRACE_HOPPER) + auto err = cudaMallocManaged(&atomic_buffer_ptr_, kAtomicBufferSize); + if (err != cudaSuccess) { + throw std::runtime_error( + std::string("cudaMallocManaged atomic buffer failed: ") + + cudaGetErrorString(err)); + } atomic_buffer_is_host_allocated_ = false; #elif defined(__HIP_PLATFORM_AMD__) || defined(__HIPCC__) - hipExtMallocWithFlags(&atomic_buffer_ptr_, kAtomicBufferSize, - hipDeviceMallocUncached); + auto err = hipExtMallocWithFlags(&atomic_buffer_ptr_, kAtomicBufferSize, + hipDeviceMallocUncached); + if (err != cudaSuccess) { + throw std::runtime_error( + std::string("hipExtMallocWithFlags atomic buffer failed: ") + + cudaGetErrorString(err)); + } atomic_buffer_is_host_allocated_ = false; #elif defined(EFA) // EFA: atomic buffer is always pinned host memory (cudaHostAlloc). - cudaHostAlloc(&atomic_buffer_ptr_, kAtomicBufferSize, - cudaHostAllocMapped | cudaHostAllocWriteCombined); + auto err = cudaHostAlloc(&atomic_buffer_ptr_, kAtomicBufferSize, + cudaHostAllocMapped | cudaHostAllocWriteCombined); + if (err != cudaSuccess) { + throw std::runtime_error(std::string("cudaHostAlloc atomic buffer failed: ") + + cudaGetErrorString(err)); + } atomic_buffer_is_host_allocated_ = true; #else // Dynamically detect: on some nodes (e.g. GH10) ibv_reg_mr fails for // cudaMalloc; use pinned host memory then. Override with // UCCL_ATOMICS_USE_HOST_MEMORY=1 to force host memory. if (can_register_gpu_memory_for_atomics(local_rank)) { - cudaMalloc(&atomic_buffer_ptr_, kAtomicBufferSize); + auto err = cudaMalloc(&atomic_buffer_ptr_, kAtomicBufferSize); + if (err != cudaSuccess) { + throw std::runtime_error( + std::string("cudaMalloc atomic buffer failed: ") + + cudaGetErrorString(err)); + } atomic_buffer_is_host_allocated_ = false; } else { - cudaHostAlloc(&atomic_buffer_ptr_, kAtomicBufferSize, - cudaHostAllocMapped); + auto err = cudaHostAlloc(&atomic_buffer_ptr_, kAtomicBufferSize, + cudaHostAllocMapped); + if (err != cudaSuccess) { + throw std::runtime_error( + std::string("cudaHostAlloc atomic buffer failed: ") + + cudaGetErrorString(err)); + } atomic_buffer_is_host_allocated_ = true; } #endif - cudaMemset(atomic_buffer_ptr_, 0, kAtomicBufferSize); + if (atomic_buffer_is_host_allocated_) { + std::memset(atomic_buffer_ptr_, 0, kAtomicBufferSize); + } else { + auto err = cudaMemset(atomic_buffer_ptr_, 0, kAtomicBufferSize); + if (err != cudaSuccess) { + throw std::runtime_error(std::string("cudaMemset atomic buffer failed: ") + + cudaGetErrorString(err)); + } + } proxy_->set_atomic_buffer_ptr(atomic_buffer_ptr_); } } @@ -140,11 +182,10 @@ void UcclProxy::stop() { throw std::runtime_error("Proxy already stopped"); } proxy_->set_progress_run(false); + proxy_->notify_proxy_thread_adaptive_sleeper(); if (thread_.joinable()) thread_.join(); running_.store(false, std::memory_order_release); - // Because proxies share the gpu_buffer, only destroy gpu_buffer for the first - // proxy. - proxy_->destroy(thread_idx_ == 0); + proxy_->destroy(false); } void UcclProxy::start(Mode m) {