diff --git a/examples/ccl/all_gather.cc b/examples/ccl/all_gather.cc index 9f640d7..c3a41b1 100644 --- a/examples/ccl/all_gather.cc +++ b/examples/ccl/all_gather.cc @@ -1,10 +1,12 @@ /** - * InfiniCCL Example: Thread-per-GPU Single-Node AllGather + * InfiniCCL Example: Thread-per-GPU Single-Node `AllGather` * - * This example validates out-of-place and in-place AllGather across two GPUs + * This example validates out-of-place and in-place `AllGather` across two GPUs * through InfiniCCL's native CCL backend without an MPI launcher. */ +#include + #include #include #include @@ -22,10 +24,9 @@ using namespace infini::ccl; namespace { -constexpr int kRankCount = 2; - struct ThreadArgs { int rank; + int size; infinicclUniqueId id; size_t num_elements; int warmup_iter; @@ -33,11 +34,11 @@ struct ThreadArgs { std::atomic_bool* all_correct; }; -bool Validate(const std::vector& output, int rank) { +bool Validate(const std::vector& output, int size, int rank) { // Every rank checks the complete gathered output by source-rank block. - const size_t num_elements = output.size() / kRankCount; + const size_t num_elements = output.size() / size; bool correct = true; - for (int source = 0; source < kRankCount; ++source) { + for (int source = 0; source < size; ++source) { const float expected = static_cast(source + 1); const size_t offset = static_cast(source) * num_elements; @@ -48,12 +49,12 @@ bool Validate(const std::vector& output, int rank) { return correct; } -void PrintResult(const std::vector& output, const char* mode, +void PrintResult(const std::vector& output, const char* mode, int size, bool correct) { const char* green = "\033[32m"; const char* red = "\033[31m"; const char* reset = "\033[0m"; - const size_t num_elements = output.size() / kRankCount; + const size_t num_elements = output.size() / size; std::cout << "\n=== " << mode << " AllGather Results ===" << std::endl; std::cout << "Correct: " @@ -61,7 +62,7 @@ void PrintResult(const std::vector& output, const char* mode, : (red + std::string("NO") + reset)) << std::endl; std::cout << "Sample blocks: "; - for (int source = 0; source < kRankCount; ++source) { + for (int source = 0; source < size; ++source) { const size_t offset = static_cast(source) * num_elements; std::cout << "[r" << source << ": " << output[offset] << "] "; } @@ -76,17 +77,17 @@ void WorkerThread(ThreadArgs args) { CHECK_RT(Rt, Rt::SetDevice(args.rank)); infinicclComm_t comm = nullptr; - CHECK_INFINI(infinicclCommInitRank(&comm, kRankCount, args.id, args.rank)); + CHECK_INFINI(infinicclCommInitRank(&comm, args.size, args.id, args.rank)); std::vector host_send(args.num_elements, static_cast(args.rank + 1)); - std::vector host_recv(args.num_elements * kRankCount, 0.0f); + std::vector host_recv(args.num_elements * args.size, 0.0f); // Prepare separate send and receive buffers for the out-of-place case. float* device_send = nullptr; float* device_recv = nullptr; const size_t send_bytes = args.num_elements * sizeof(float); - const size_t recv_bytes = send_bytes * kRankCount; + const size_t recv_bytes = send_bytes * args.size; CHECK_RT(Rt, Rt::Malloc(reinterpret_cast(&device_send), send_bytes)); CHECK_RT(Rt, Rt::Malloc(reinterpret_cast(&device_recv), recv_bytes)); @@ -115,12 +116,12 @@ void WorkerThread(ThreadArgs args) { CHECK_RT(Rt, Rt::Memcpy(host_recv.data(), device_recv, recv_bytes, Rt::MemcpyDeviceToHost)); - const bool out_of_place_correct = Validate(host_recv, args.rank); + const bool out_of_place_correct = Validate(host_recv, args.size, args.rank); if (!out_of_place_correct) { args.all_correct->store(false, std::memory_order_relaxed); } if (args.rank == 0) { - PrintResult(host_recv, "Out-of-place", out_of_place_correct); + PrintResult(host_recv, "Out-of-place", args.size, out_of_place_correct); } // Seed only the local block, then use it as both input and output. @@ -138,16 +139,16 @@ void WorkerThread(ThreadArgs args) { CHECK_RT(Rt, Rt::StreamSynchronize(nullptr)); CHECK_RT(Rt, Rt::Memcpy(host_recv.data(), device_recv, recv_bytes, Rt::MemcpyDeviceToHost)); - const bool in_place_correct = Validate(host_recv, args.rank); + const bool in_place_correct = Validate(host_recv, args.size, args.rank); if (!in_place_correct) { args.all_correct->store(false, std::memory_order_relaxed); } if (args.rank == 0) { - PrintResult(host_recv, "In-place", in_place_correct); + PrintResult(host_recv, "In-place", args.size, in_place_correct); std::cout << "\n=== Single-Node Threaded AllGather Results ===" << std::endl; - Metrics metrics{elapsed, recv_bytes, kRankCount}; + Metrics metrics{elapsed, recv_bytes, args.size}; metrics.Print(); } @@ -159,25 +160,57 @@ void WorkerThread(ThreadArgs args) { } // namespace -int main() { - constexpr size_t kNumElements = 1 << 20; - constexpr int kWarmupIterations = 2; - constexpr int kProfileIterations = 20; +int main(int argc, char** argv) { + int num_gpus = 8; + int warmup_iters = 2; + int profile_iters = 20; + size_t num_elements = 1 << 20; + + int opt; + while ((opt = getopt(argc, argv, "g:w:p:n:h")) != -1) { + switch (opt) { + case 'g': + num_gpus = std::stoi(optarg); + break; + case 'w': + warmup_iters = std::stoi(optarg); + break; + case 'p': + profile_iters = std::stoi(optarg); + break; + case 'n': + num_elements = static_cast(std::stoull(optarg)); + break; + case 'h': + std::cout << "Usage: " << argv[0] << " [options]\n" + << "Options:\n" + << " -g Number of GPUs (default: 8)\n" + << " -w Warmup iterations (default: 2)\n" + << " -p Profile iterations (default: 20)\n" + << " -n Number of elements (default: " + << (1 << 20) << ")\n"; + return EXIT_SUCCESS; + default: + std::cerr << "Invalid argument. Use -h for help." << std::endl; + return EXIT_FAILURE; + } + } + + char hostname[256]; + gethostname(hostname, sizeof(hostname)); + std::cout << "[Main Process] Host: " << hostname + << " | Target GPUs: " << num_gpus << std::endl; infinicclUniqueId shared_id; CHECK_INFINI(infinicclGetUniqueId(&shared_id)); std::atomic_bool all_correct{true}; std::vector threads; - threads.reserve(kRankCount); - - for (int rank = 0; rank < kRankCount; ++rank) { - ThreadArgs args{rank, - shared_id, - kNumElements, - kWarmupIterations, - kProfileIterations, - &all_correct}; + threads.reserve(num_gpus); + + for (int rank = 0; rank < num_gpus; ++rank) { + ThreadArgs args{rank, num_gpus, shared_id, num_elements, + warmup_iters, profile_iters, &all_correct}; threads.emplace_back(WorkerThread, args); } for (auto& thread : threads) {