use std::ffi::c_char;
use std::mem::MaybeUninit;
type NcclByte = c_char;
use anyhow::{Context, Result};
use cudarc::driver::sys::CUstream;
use cudarc::nccl::sys::{
ncclComm_t, ncclCommInitRank, ncclGetUniqueId, ncclResult_t, ncclUniqueId,
};
#[derive(Clone)]
pub struct NcclBootstrap {
nccl_id: ncclUniqueId,
world_size: usize,
}
impl NcclBootstrap {
pub fn generate(world_size: usize) -> Result<Self> {
anyhow::ensure!(
world_size > 0 && world_size <= i32::MAX as usize,
"world_size must be in 1..={}, got {}",
i32::MAX,
world_size
);
let mut nccl_id = MaybeUninit::<ncclUniqueId>::uninit();
let result = unsafe { ncclGetUniqueId(nccl_id.as_mut_ptr()) };
check_nccl_result(result).context("Failed to generate NCCL unique ID")?;
let nccl_id = unsafe { nccl_id.assume_init() };
Ok(Self {
nccl_id,
world_size,
})
}
pub fn world_size(&self) -> usize {
self.world_size
}
pub fn serialize(&self) -> Vec<u8> {
let mut bytes = Vec::with_capacity(8 + 128);
bytes.extend_from_slice(&(self.world_size as u64).to_le_bytes());
for &byte in &self.nccl_id.internal {
bytes.push(byte as u8);
}
bytes
}
pub fn deserialize(bytes: &[u8]) -> Result<Self> {
if bytes.len() != 8 + 128 {
anyhow::bail!(
"Invalid bootstrap data length: expected {}, got {}",
8 + 128,
bytes.len()
);
}
let world_size = u64::from_le_bytes(bytes[0..8].try_into().unwrap()) as usize;
let mut nccl_id = ncclUniqueId {
internal: [0 as NcclByte; 128],
};
for (i, &byte) in bytes[8..].iter().enumerate() {
nccl_id.internal[i] = byte as NcclByte;
}
Ok(Self {
nccl_id,
world_size,
})
}
pub fn init_communicator(&self, rank: usize, _stream: CUstream) -> Result<ncclComm_t> {
if rank >= self.world_size {
anyhow::bail!(
"Rank {} is invalid for world_size {}",
rank,
self.world_size
);
}
anyhow::ensure!(
self.world_size <= i32::MAX as usize,
"world_size {} exceeds i32::MAX",
self.world_size
);
let mut comm = MaybeUninit::<ncclComm_t>::uninit();
let result = unsafe {
ncclCommInitRank(
comm.as_mut_ptr(),
self.world_size as i32,
self.nccl_id,
rank as i32,
)
};
check_nccl_result(result).context("Failed to initialize NCCL communicator")?;
let comm = unsafe { comm.assume_init() };
tracing::debug!(
rank,
world_size = self.world_size,
"NCCL communicator initialized"
);
Ok(comm)
}
}
pub(crate) fn check_nccl_result(result: ncclResult_t) -> Result<()> {
if result == ncclResult_t::ncclSuccess {
Ok(())
} else {
anyhow::bail!("NCCL error: {:?}", result)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_bootstrap_serialization_roundtrip() {
let world_size = 4;
let original = NcclBootstrap {
nccl_id: ncclUniqueId {
internal: [42 as NcclByte; 128],
},
world_size,
};
let bytes = original.serialize();
assert_eq!(bytes.len(), 8 + 128);
let deserialized = NcclBootstrap::deserialize(&bytes).unwrap();
assert_eq!(deserialized.world_size, world_size);
assert_eq!(deserialized.nccl_id.internal, original.nccl_id.internal);
}
#[test]
fn test_deserialize_invalid_length() {
let bytes = vec![0u8; 10]; let result = NcclBootstrap::deserialize(&bytes);
assert!(result.is_err());
}
}