use std::path::Path;
use std::time::{Duration, Instant};
use cudarc::nccl::{result, sys};
use super::error::DistError;
const MIN_NCCL_VERSION: i32 = 22600;
pub struct MambaComm {
comm: sys::ncclComm_t,
rank: usize,
world: usize,
aborted: std::sync::Arc<std::sync::atomic::AtomicBool>,
_cuda_ctx: std::sync::Arc<cudarc::driver::CudaContext>,
}
unsafe impl Send for MambaComm {}
impl MambaComm {
pub fn preflight_version() -> Result<i32, DistError> {
let v = result::get_nccl_version()
.map_err(|e| DistError::Transport(format!("NCCL version query: {e:?}")))?;
if v < MIN_NCCL_VERSION {
return Err(DistError::Transport(format!(
"NCCL {v} is below the supported floor {MIN_NCCL_VERSION}"
)));
}
Ok(v)
}
pub fn exchange_unique_id(
path: &Path,
rank: usize,
timeout: Duration,
) -> Result<sys::ncclUniqueId, DistError> {
if rank == 0 {
let id = result::get_uniqueid()
.map_err(|e| DistError::Transport(format!("NCCL unique id: {e:?}")))?;
let bytes: Vec<u8> = id.internal.iter().map(|&b| b as u8).collect();
let tmp = path.with_extension("tmp");
std::fs::write(&tmp, &bytes)
.map_err(|e| DistError::Rendezvous(format!("write {}: {e}", tmp.display())))?;
std::fs::rename(&tmp, path)
.map_err(|e| DistError::Rendezvous(format!("rename {}: {e}", path.display())))?;
Ok(id)
} else {
let deadline = Instant::now() + timeout;
loop {
if let Ok(bytes) = std::fs::read(path)
&& bytes.len() == 128
{
let mut id = sys::ncclUniqueId { internal: [0; 128] };
for (dst, &src) in id.internal.iter_mut().zip(bytes.iter()) {
*dst = src as core::ffi::c_char;
}
return Ok(id);
}
if Instant::now() >= deadline {
return Err(DistError::Rendezvous(format!(
"rank {rank}: NCCL unique id did not appear at {} within {timeout:?}",
path.display()
)));
}
std::thread::sleep(Duration::from_millis(5));
}
}
}
pub fn init_with_deadline(
unique_id: sys::ncclUniqueId,
rank: usize,
world: usize,
cuda_ctx: std::sync::Arc<cudarc::driver::CudaContext>,
deadline: Duration,
) -> Result<Self, DistError> {
super::watchdog::run_with_deadline("nccl-init", deadline, move || {
cuda_ctx
.bind_to_thread()
.map_err(|e| DistError::Transport(format!("bind CUDA ctx for init: {e:?}")))?;
Self::init(unique_id, rank, world, cuda_ctx)
})?
}
pub(super) fn with_watchdog<R>(
&self,
name: &str,
deadline: Duration,
f: impl FnOnce() -> Result<R, DistError>,
) -> Result<R, DistError> {
let comm_addr = self.comm as usize;
let aborted = self.aborted.clone();
let wd = super::watchdog::Watchdog::arm(name, deadline, move || {
aborted.store(true, std::sync::atomic::Ordering::Release);
let _ = unsafe { result::comm_abort(comm_addr as sys::ncclComm_t) };
})?;
let out = f();
if wd.disarm() {
return Err(DistError::Transport(format!(
"{name}: collective window exceeded {deadline:?} — communicator aborted"
)));
}
out
}
pub fn init(
unique_id: sys::ncclUniqueId,
rank: usize,
world: usize,
cuda_ctx: std::sync::Arc<cudarc::driver::CudaContext>,
) -> Result<Self, DistError> {
let mut comm: sys::ncclComm_t = std::ptr::null_mut();
unsafe { result::comm_init_rank(&mut comm, world as i32, unique_id, rank as i32) }
.map_err(|e| {
let detail = unsafe {
let p = sys::ncclGetLastError(std::ptr::null_mut());
if p.is_null() {
String::new()
} else {
std::ffi::CStr::from_ptr(p).to_string_lossy().into_owned()
}
};
DistError::Transport(format!("NCCL init rank {rank}/{world}: {e:?} — {detail}"))
})?;
Ok(Self {
comm,
rank,
world,
aborted: std::sync::Arc::new(std::sync::atomic::AtomicBool::new(false)),
_cuda_ctx: cuda_ctx,
})
}
pub fn rank(&self) -> usize {
self.rank
}
pub fn world(&self) -> usize {
self.world
}
pub fn all_reduce_sum_f32(
&self,
ptr: cudarc::driver::sys::CUdeviceptr,
count: usize,
stream: &cudarc::driver::CudaStream,
) -> Result<(), DistError> {
unsafe {
result::all_reduce(
ptr as *const core::ffi::c_void,
ptr as *mut core::ffi::c_void,
count,
sys::ncclDataType_t::ncclFloat32,
sys::ncclRedOp_t::ncclSum,
self.comm,
stream.cu_stream() as *mut _,
)
}
.map_err(|e| DistError::Transport(format!("NCCL allreduce({count} f32): {e:?}")))?;
Ok(())
}
pub(super) fn send_f32(
&self,
ptr: cudarc::driver::sys::CUdeviceptr,
count: usize,
peer: usize,
stream: &cudarc::driver::CudaStream,
) -> Result<(), DistError> {
unsafe {
result::send(
ptr as *const core::ffi::c_void,
count,
sys::ncclDataType_t::ncclFloat32,
peer as core::ffi::c_int,
self.comm,
stream.cu_stream() as *mut _,
)
}
.map_err(|e| DistError::Transport(format!("NCCL send({count} f32 -> {peer}): {e:?}")))?;
Ok(())
}
pub(super) fn recv_f32(
&self,
ptr: cudarc::driver::sys::CUdeviceptr,
count: usize,
peer: usize,
stream: &cudarc::driver::CudaStream,
) -> Result<(), DistError> {
unsafe {
result::recv(
ptr as *mut core::ffi::c_void,
count,
sys::ncclDataType_t::ncclFloat32,
peer as core::ffi::c_int,
self.comm,
stream.cu_stream() as *mut _,
)
}
.map_err(|e| DistError::Transport(format!("NCCL recv({count} f32 <- {peer}): {e:?}")))?;
Ok(())
}
pub(super) fn broadcast_f32(
&self,
ptr: cudarc::driver::sys::CUdeviceptr,
count: usize,
root: usize,
stream: &cudarc::driver::CudaStream,
) -> Result<(), DistError> {
unsafe {
result::broadcast(
ptr as *const core::ffi::c_void,
ptr as *mut core::ffi::c_void,
count,
sys::ncclDataType_t::ncclFloat32,
root as core::ffi::c_int,
self.comm,
stream.cu_stream() as *mut _,
)
}
.map_err(|e| {
DistError::Transport(format!("NCCL broadcast({count} f32, root {root}): {e:?}"))
})?;
Ok(())
}
pub(super) fn group<R>(f: impl FnOnce() -> Result<R, DistError>) -> Result<R, DistError> {
result::group_start()
.map_err(|e| DistError::Transport(format!("NCCL group start: {e:?}")))?;
let out = f();
let end = result::group_end();
let r = out?;
end.map_err(|e| DistError::Transport(format!("NCCL group end: {e:?}")))?;
Ok(r)
}
pub fn all_reduce_max_i32(
&self,
ptr: cudarc::driver::sys::CUdeviceptr,
count: usize,
stream: &cudarc::driver::CudaStream,
) -> Result<(), DistError> {
unsafe {
result::all_reduce(
ptr as *const core::ffi::c_void,
ptr as *mut core::ffi::c_void,
count,
sys::ncclDataType_t::ncclInt32,
sys::ncclRedOp_t::ncclMax,
self.comm,
stream.cu_stream() as *mut _,
)
}
.map_err(|e| DistError::Transport(format!("NCCL allreduce max({count} i32): {e:?}")))?;
Ok(())
}
pub fn shutdown(mut self) -> Result<(), DistError> {
let comm = std::mem::replace(&mut self.comm, std::ptr::null_mut());
if self.aborted.load(std::sync::atomic::Ordering::Acquire) {
return Err(DistError::Transport(
"communicator was aborted by the collective watchdog — teardown already complete"
.into(),
));
}
unsafe {
if let Err(e) = result::comm_finalize(comm) {
let _ = result::comm_abort(comm);
return Err(DistError::Transport(format!("NCCL finalize: {e:?}")));
}
if let Err(e) = result::comm_destroy(comm) {
return Err(DistError::Transport(format!("NCCL destroy: {e:?}")));
}
}
Ok(())
}
}
impl Drop for MambaComm {
fn drop(&mut self) {
if !self.comm.is_null() && !self.aborted.load(std::sync::atomic::Ordering::Acquire) {
let _ = unsafe { result::comm_abort(self.comm) };
}
}
}