Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
272 changes: 272 additions & 0 deletions examples/ccl/broadcast.cc
Original file line number Diff line number Diff line change
@@ -0,0 +1,272 @@
/**
* InfiniCCL Example: Thread-per-GPU Single-Node Broadcast
*
* This example spawns one CPU thread per GPU and performs native CCL
* broadcasts without an MPI launcher.
*/

#include <unistd.h>

#include <algorithm>
#include <array>
#include <atomic>
#include <charconv>
#include <cmath>
#include <cstdlib>
#include <iomanip>
#include <iostream>
#include <limits>
#include <string>
#include <string_view>
#include <thread>
#include <vector>

#include "backend_manifest.h"
#include "infiniccl.h"
#include "utils.h"

using namespace infini::ccl;

struct ThreadArgs {
int rank;
int size;
int root;
infinicclUniqueId id;
size_t num_elements;
int warmup_iter;
int profile_iter;
std::atomic_bool *all_correct;
};

template <typename T>
bool ParseIntegerOption(const char *argument, T *value) {
if (!argument || !value) {
return false;
}

const std::string_view text(argument);
if (text.empty()) {
return false;
}

T parsed{};
const auto [end, error] =
std::from_chars(text.data(), text.data() + text.size(), parsed);
if (error != std::errc{} || end != text.data() + text.size()) {
return false;
}

*value = parsed;
return true;
}

void PrintBroadcastMetrics(size_t num_elements, double elapsed_ms) {
constexpr double kBytesPerMiB = 1024.0 * 1024.0;
constexpr double kBytesPerGB = 1.0e9;
const double data_bytes = static_cast<double>(num_elements) * sizeof(float);
const auto original_flags = std::cout.flags();
const auto original_precision = std::cout.precision();

std::cout << std::fixed;
std::cout << "Data size: " << std::setprecision(2)
<< data_bytes / kBytesPerMiB << " MiB" << std::endl;
std::cout << "Time: " << std::setprecision(3) << elapsed_ms << " ms"
<< std::endl;
if (elapsed_ms > 0.0 && std::isfinite(elapsed_ms)) {
const double algorithm_bandwidth =
data_bytes / kBytesPerGB / (elapsed_ms / 1000.0);
const double bus_bandwidth = algorithm_bandwidth;
std::cout << "Throughput: " << std::setprecision(2) << bus_bandwidth
<< " GB/s (Bus BW)" << std::endl;
std::cout << "Alg Bandwidth: " << algorithm_bandwidth << " GB/s"
<< std::endl;
} else {
std::cout << "Throughput: N/A (Bus BW)" << std::endl;
std::cout << "Alg Bandwidth: N/A" << std::endl;
}

std::cout.flags(original_flags);
std::cout.precision(original_precision);
}

void WorkerThread(ThreadArgs args) {
constexpr Device::Type kDevType =
ListGetBest<DevicePriority>(EnabledDevices{});
using Rt = Runtime<kDevType>;

CHECK_RT(Rt, Rt::SetDevice(args.rank));

infinicclComm_t comm = nullptr;
CHECK_INFINI(infinicclCommInitRank(&comm, args.size, args.id, args.rank));

constexpr float kRootValue = 42.0f;
constexpr float kSentinelValue = -1.0f;
const size_t total_bytes = args.num_elements * sizeof(float);

std::vector<float> h_send(args.num_elements, kSentinelValue);
std::vector<float> h_recv(args.num_elements, kSentinelValue);

float *d_send = nullptr;
float *d_recv = nullptr;
CHECK_RT(Rt, Rt::Malloc(reinterpret_cast<void **>(&d_send), total_bytes));
CHECK_RT(Rt, Rt::Malloc(reinterpret_cast<void **>(&d_recv), total_bytes));

auto ResetBuffers = [&]() {
std::fill(h_send.begin(), h_send.end(),
args.rank == args.root ? kRootValue : kSentinelValue);
std::fill(h_recv.begin(), h_recv.end(), kSentinelValue);
CHECK_RT(Rt, Rt::Memcpy(d_send, h_send.data(), total_bytes,
Rt::MemcpyHostToDevice));
CHECK_RT(Rt, Rt::Memcpy(d_recv, h_recv.data(), total_bytes,
Rt::MemcpyHostToDevice));
CHECK_RT(Rt, Rt::StreamSynchronize(nullptr));
};

auto RunScenario = [&](const std::string &name, float *verify_buff,
auto &&collective_call) {
for (int i = 0; i < args.warmup_iter; ++i) {
CHECK_INFINI(collective_call());
}
CHECK_RT(Rt, Rt::StreamSynchronize(nullptr));

Timer timer;
for (int i = 0; i < args.profile_iter; ++i) {
CHECK_INFINI(collective_call());
}
CHECK_RT(Rt, Rt::StreamSynchronize(nullptr));
const double elapsed =
timer.ElapsedMs() / static_cast<double>(args.profile_iter);

CHECK_RT(Rt, Rt::Memcpy(h_recv.data(), verify_buff, total_bytes,
Rt::MemcpyDeviceToHost));

const bool correct = Validator::ValidateResult(
h_recv.data(), args.num_elements, kRootValue, args.rank, true, name);
if (!correct) {
args.all_correct->store(false, std::memory_order_relaxed);
}

if (args.rank == 0) {
PrintBroadcastMetrics(args.num_elements, elapsed);
}
};

ResetBuffers();
RunScenario("Out-of-Place Broadcast", d_recv, [&]() {
return infinicclBroadcast(d_send, d_recv, args.num_elements,
infinicclFloat32, args.root, comm, nullptr);
});

ResetBuffers();
float *d_in_place = args.rank == args.root ? d_send : d_recv;
RunScenario("In-Place Broadcast", d_in_place, [&]() {
return infinicclBroadcast(d_in_place, d_in_place, args.num_elements,
infinicclFloat32, args.root, comm, nullptr);
});

ResetBuffers();
d_in_place = args.rank == args.root ? d_send : d_recv;
RunScenario("Legacy In-Place Bcast", d_in_place, [&]() {
return infinicclBcast(d_in_place, args.num_elements, infinicclFloat32,
args.root, comm, nullptr);
});

CHECK_RT(Rt, Rt::Free(d_send));
CHECK_RT(Rt, Rt::Free(d_recv));
CHECK_INFINI(infinicclCommDestroy(comm));
}

int main(int argc, char **argv) {
int num_gpus = 8;
int warmup_iters = 1;
int profile_iters = 20;
size_t num_elements = 1 << 25;

int opt;
while ((opt = getopt(argc, argv, "g:w:p:n:h")) != -1) {
switch (opt) {
case 'g':
if (!ParseIntegerOption(optarg, &num_gpus)) {
std::cerr << "Invalid value for `-g`." << std::endl;
return EXIT_FAILURE;
}
break;
case 'w':
if (!ParseIntegerOption(optarg, &warmup_iters)) {
std::cerr << "Invalid value for `-w`." << std::endl;
return EXIT_FAILURE;
}
break;
case 'p':
if (!ParseIntegerOption(optarg, &profile_iters)) {
std::cerr << "Invalid value for `-p`." << std::endl;
return EXIT_FAILURE;
}
break;
case 'n':
if (!ParseIntegerOption(optarg, &num_elements)) {
std::cerr << "Invalid value for `-n`." << std::endl;
return EXIT_FAILURE;
}
break;
case 'h':
std::cout << "Usage: " << argv[0] << " [options]\n"
<< "Options:\n"
<< " -g <num_gpus> Number of GPUs (default: 8)\n"
<< " -w <warmup_iters> Warmup iterations (default: 1)\n"
<< " -p <profile_iters> Profile iterations (default: 20)\n"
<< " -n <num_elements> Number of elements (default: "
<< (1 << 25) << ")\n";
return EXIT_SUCCESS;
default:
std::cerr << "Invalid argument. Use `-h` for help." << std::endl;
return EXIT_FAILURE;
}
}

if (num_gpus <= 0 || warmup_iters < 0 || profile_iters <= 0 ||
num_elements == 0 ||
num_elements > std::numeric_limits<size_t>::max() / sizeof(float)) {
std::cerr << "Invalid execution parameter." << std::endl;
return EXIT_FAILURE;
}

std::array<char, 256> hostname{};
if (gethostname(hostname.data(), hostname.size()) != 0) {
std::cerr << "Failed to query the host name." << std::endl;
return EXIT_FAILURE;
}
hostname.back() = '\0';

const int root = num_gpus > 1 ? num_gpus - 1 : 0;
std::cout << "[Main Process] Host: " << hostname.data()
<< " | Target GPUs: " << num_gpus << " | Root: " << root
<< std::endl;

infinicclUniqueId shared_id;
CHECK_INFINI(infinicclGetUniqueId(&shared_id));

std::atomic_bool all_correct{true};
std::vector<std::thread> threads;
threads.reserve(num_gpus);

for (int rank = 0; rank < num_gpus; ++rank) {
ThreadArgs args{rank, num_gpus, root, shared_id,
num_elements, warmup_iters, profile_iters, &all_correct};
threads.emplace_back(WorkerThread, args);
}

for (auto &thread : threads) {
if (thread.joinable()) {
thread.join();
}
}

if (!all_correct.load(std::memory_order_relaxed)) {
std::cerr << "Broadcast validation failed." << std::endl;
return EXIT_FAILURE;
}

std::cout << "[Main Process] All broadcast scenarios passed." << std::endl;
return EXIT_SUCCESS;
}
Loading
Loading