use std::ffi::c_void;
use std::ptr;
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::Arc;
use flodl_sys::{self as ffi, FlodlTensor};
use crate::tensor::{
check_err, current_cuda_device, set_current_cuda_device,
Device, Result, Tensor, TensorError,
};
use crate::tensor::cuda_stream::CudaStream;
#[derive(Clone, Copy, Debug)]
#[repr(i32)]
pub enum ReduceOp {
Sum = 0,
Prod = 1,
Max = 2,
Min = 3,
Avg = 4,
}
pub struct NcclComms {
handle: *mut c_void,
devices: Vec<Device>,
}
unsafe impl Send for NcclComms {}
impl NcclComms {
pub fn new(devices: &[Device]) -> Result<Self> {
if devices.len() < 2 {
return Err(TensorError::new(
"NcclComms requires at least 2 devices",
));
}
let mut devlist: Vec<i32> = Vec::with_capacity(devices.len());
for &dev in devices {
match dev {
Device::CUDA(idx) => devlist.push(idx as i32),
Device::CPU => {
return Err(TensorError::new(
"NcclComms requires CUDA devices, got CPU",
))
}
}
}
let mut handle: *mut c_void = ptr::null_mut();
let saved = current_cuda_device();
let err = unsafe {
ffi::flodl_nccl_init(
devlist.len() as i32,
devlist.as_ptr(),
&mut handle,
)
};
set_current_cuda_device(saved);
check_err(err)?;
Ok(NcclComms {
handle,
devices: devices.to_vec(),
})
}
pub fn all_reduce(&self, tensors: &[&Tensor], op: ReduceOp) -> Result<()> {
self.validate_tensors(tensors, "all_reduce")?;
let mut handles: Vec<FlodlTensor> = tensors.iter().map(|t| t.handle).collect();
let saved = current_cuda_device();
let err = unsafe {
ffi::flodl_nccl_all_reduce(
self.handle,
handles.as_mut_ptr(),
ptr::null_mut(),
op as i32,
)
};
set_current_cuda_device(saved);
check_err(err)
}
pub fn all_reduce_on_streams(
&self,
tensors: &[&Tensor],
op: ReduceOp,
streams: &[&CudaStream],
) -> Result<()> {
self.validate_tensors(tensors, "all_reduce_on_streams")?;
if streams.len() != self.devices.len() {
return Err(TensorError::new(&format!(
"all_reduce_on_streams: expected {} streams, got {}",
self.devices.len(), streams.len()
)));
}
let mut handles: Vec<FlodlTensor> = tensors.iter().map(|t| t.handle).collect();
let mut stream_ptrs: Vec<*mut c_void> = streams.iter().map(|s| s.as_ptr()).collect();
let saved = current_cuda_device();
let err = unsafe {
ffi::flodl_nccl_all_reduce(
self.handle,
handles.as_mut_ptr(),
stream_ptrs.as_mut_ptr(),
op as i32,
)
};
set_current_cuda_device(saved);
check_err(err)
}
pub fn broadcast(&self, tensors: &[&Tensor], root: usize) -> Result<()> {
self.validate_tensors(tensors, "broadcast")?;
if root >= self.devices.len() {
return Err(TensorError::new(&format!(
"broadcast: root {} out of range (have {} devices)",
root, self.devices.len()
)));
}
let mut handles: Vec<FlodlTensor> = tensors.iter().map(|t| t.handle).collect();
let saved = current_cuda_device();
let err = unsafe {
ffi::flodl_nccl_broadcast(
self.handle,
handles.as_mut_ptr(),
ptr::null_mut(),
root as i32,
)
};
set_current_cuda_device(saved);
check_err(err)
}
pub fn broadcast_on_streams(
&self,
tensors: &[&Tensor],
root: usize,
streams: &[&CudaStream],
) -> Result<()> {
self.validate_tensors(tensors, "broadcast_on_streams")?;
if root >= self.devices.len() {
return Err(TensorError::new(&format!(
"broadcast_on_streams: root {} out of range", root
)));
}
if streams.len() != self.devices.len() {
return Err(TensorError::new(&format!(
"broadcast_on_streams: expected {} streams, got {}",
self.devices.len(), streams.len()
)));
}
let mut handles: Vec<FlodlTensor> = tensors.iter().map(|t| t.handle).collect();
let mut stream_ptrs: Vec<*mut c_void> = streams.iter().map(|s| s.as_ptr()).collect();
let saved = current_cuda_device();
let err = unsafe {
ffi::flodl_nccl_broadcast(
self.handle,
handles.as_mut_ptr(),
stream_ptrs.as_mut_ptr(),
root as i32,
)
};
set_current_cuda_device(saved);
check_err(err)
}
pub fn size(&self) -> usize {
self.devices.len()
}
pub fn devices(&self) -> &[Device] {
&self.devices
}
fn validate_tensors(&self, tensors: &[&Tensor], op: &str) -> Result<()> {
if tensors.len() != self.devices.len() {
return Err(TensorError::new(&format!(
"{}: expected {} tensors (one per device), got {}",
op, self.devices.len(), tensors.len()
)));
}
Ok(())
}
pub fn split(self) -> Result<Vec<NcclRankComm>> {
let mut comms = Vec::with_capacity(self.devices.len());
for i in 0..self.devices.len() {
let mut rank_handle: *mut c_void = ptr::null_mut();
let err = unsafe {
ffi::flodl_nccl_split_rank(
self.handle,
i as i32,
&mut rank_handle,
)
};
check_err(err)?;
let abort_handle = Arc::new(NcclAbortHandle {
ptr: rank_handle,
aborted: AtomicBool::new(false),
guard: std::sync::Mutex::new(()),
});
comms.push(NcclRankComm {
handle: rank_handle,
rank: i,
world_size: self.devices.len(),
abort_handle,
});
}
Ok(comms)
}
}
impl Drop for NcclComms {
fn drop(&mut self) {
if !self.handle.is_null() {
unsafe { ffi::flodl_nccl_destroy(self.handle) };
self.handle = ptr::null_mut();
}
}
}
pub const NCCL_UNIQUE_ID_BYTES: usize = 128;
#[derive(Clone)]
pub struct NcclUniqueId {
bytes: [u8; NCCL_UNIQUE_ID_BYTES],
}
unsafe impl Send for NcclUniqueId {}
unsafe impl Sync for NcclUniqueId {}
impl NcclUniqueId {
pub fn new() -> Result<Self> {
let mut bytes = [0u8; NCCL_UNIQUE_ID_BYTES];
let err = unsafe { ffi::flodl_nccl_get_unique_id(bytes.as_mut_ptr()) };
check_err(err)?;
Ok(NcclUniqueId { bytes })
}
pub fn as_bytes(&self) -> &[u8; NCCL_UNIQUE_ID_BYTES] {
&self.bytes
}
pub fn from_bytes(bytes: [u8; NCCL_UNIQUE_ID_BYTES]) -> Self {
NcclUniqueId { bytes }
}
}
impl std::fmt::Debug for NcclUniqueId {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("NcclUniqueId").finish()
}
}
pub struct NcclAbortHandle {
ptr: *mut c_void,
aborted: AtomicBool,
guard: std::sync::Mutex<()>,
}
unsafe impl Send for NcclAbortHandle {}
unsafe impl Sync for NcclAbortHandle {}
impl NcclAbortHandle {
pub fn abort(&self) -> Result<()> {
let _guard = self.guard.lock().expect("nccl issue guard poisoned");
if !self.claim() {
return Ok(()); }
let err = unsafe { ffi::flodl_nccl_abort_rank(self.ptr) };
check_err(err)
}
pub(crate) fn lock_for_issue(&self) -> Result<std::sync::MutexGuard<'_, ()>> {
let guard = self.guard.lock().expect("nccl issue guard poisoned");
if self.is_aborted() {
return Err(TensorError::new(
"NCCL communicator aborted (peer death); collective refused — \
the caller must rebuild the comm on the surviving cohort",
));
}
Ok(guard)
}
pub fn is_aborted(&self) -> bool {
self.aborted.load(Ordering::Acquire)
}
fn claim(&self) -> bool {
!self.aborted.swap(true, Ordering::AcqRel)
}
}
pub struct NcclRankComm {
handle: *mut c_void,
rank: usize,
world_size: usize,
abort_handle: Arc<NcclAbortHandle>,
}
unsafe impl Send for NcclRankComm {}
impl NcclRankComm {
pub fn init_rank(rank: usize, world_size: usize, uid: &NcclUniqueId) -> Result<Self> {
if rank >= world_size {
return Err(TensorError::new(&format!(
"NcclRankComm: rank {} >= world_size {}", rank, world_size
)));
}
if world_size < 2 {
return Err(TensorError::new(
"NcclRankComm requires world_size >= 2"
));
}
let mut handle: *mut c_void = ptr::null_mut();
let err = unsafe {
ffi::flodl_nccl_init_rank(
rank as i32,
world_size as i32,
uid.bytes.as_ptr(),
&mut handle,
)
};
check_err(err)?;
let abort_handle = Arc::new(NcclAbortHandle {
ptr: handle,
aborted: AtomicBool::new(false),
guard: std::sync::Mutex::new(()),
});
Ok(NcclRankComm { handle, rank, world_size, abort_handle })
}
pub fn rank(&self) -> usize {
self.rank
}
pub fn world_size(&self) -> usize {
self.world_size
}
pub fn abort_handle(&self) -> Arc<NcclAbortHandle> {
self.abort_handle.clone()
}
pub fn all_reduce(&self, tensors: &[&Tensor], op: ReduceOp) -> Result<()> {
let _guard = self.abort_handle.lock_for_issue()?;
let mut handles: Vec<ffi::FlodlTensor> = tensors.iter().map(|t| t.handle).collect();
let err = unsafe {
ffi::flodl_nccl_all_reduce_rank(
self.handle,
handles.as_mut_ptr(),
handles.len() as i32,
ptr::null_mut(),
op as i32,
)
};
check_err(err)
}
pub fn all_reduce_premul_sum(
&self,
tensors: &[&Tensor],
factor: f32,
stream: Option<&CudaStream>,
) -> Result<()> {
for (i, t) in tensors.iter().enumerate() {
if t.dtype() != crate::tensor::DType::Float32 {
return Err(crate::TensorError::new(&format!(
"all_reduce_premul_sum: tensor {i} is {:?}, but the \
PreMulSum scalar is f32 — NCCL requires matching dtypes",
t.dtype(),
)));
}
}
let _guard = self.abort_handle.lock_for_issue()?;
let mut op: i32 = 0;
let err = unsafe {
ffi::flodl_nccl_redop_premulsum_create_rank(self.handle, factor, &mut op)
};
check_err(err)?;
let mut handles: Vec<ffi::FlodlTensor> = tensors.iter().map(|t| t.handle).collect();
let stream_ptr = stream.map_or(ptr::null_mut(), |s| s.as_ptr());
let reduce_err = unsafe {
ffi::flodl_nccl_all_reduce_rank(
self.handle,
handles.as_mut_ptr(),
handles.len() as i32,
stream_ptr,
op,
)
};
let destroy_err = unsafe { ffi::flodl_nccl_redop_destroy_rank(self.handle, op) };
check_err(reduce_err)?;
check_err(destroy_err)
}
pub fn all_reduce_on_stream(
&self,
tensors: &[&Tensor],
op: ReduceOp,
stream: &CudaStream,
) -> Result<()> {
let _guard = self.abort_handle.lock_for_issue()?;
let mut handles: Vec<ffi::FlodlTensor> = tensors.iter().map(|t| t.handle).collect();
let err = unsafe {
ffi::flodl_nccl_all_reduce_rank(
self.handle,
handles.as_mut_ptr(),
handles.len() as i32,
stream.as_ptr(),
op as i32,
)
};
check_err(err)
}
pub fn broadcast(&self, tensors: &[&Tensor], root: usize) -> Result<()> {
if root >= self.world_size {
return Err(crate::TensorError::new(&format!(
"NcclRankComm::broadcast: root {root} out of range \
(world_size = {})",
self.world_size
)));
}
let _guard = self.abort_handle.lock_for_issue()?;
let mut handles: Vec<ffi::FlodlTensor> = tensors.iter().map(|t| t.handle).collect();
let err = unsafe {
ffi::flodl_nccl_broadcast_rank(
self.handle,
handles.as_mut_ptr(),
handles.len() as i32,
ptr::null_mut(),
root as i32,
)
};
check_err(err)
}
pub fn broadcast_on_stream(
&self,
tensors: &[&Tensor],
root: usize,
stream: &CudaStream,
) -> Result<()> {
if root >= self.world_size {
return Err(crate::TensorError::new(&format!(
"NcclRankComm::broadcast_on_stream: root {root} out of range \
(world_size = {})",
self.world_size
)));
}
let _guard = self.abort_handle.lock_for_issue()?;
let mut handles: Vec<ffi::FlodlTensor> = tensors.iter().map(|t| t.handle).collect();
let err = unsafe {
ffi::flodl_nccl_broadcast_rank(
self.handle,
handles.as_mut_ptr(),
handles.len() as i32,
stream.as_ptr(),
root as i32,
)
};
check_err(err)
}
}
impl Drop for NcclRankComm {
fn drop(&mut self) {
if self.abort_handle.claim() && !self.handle.is_null() {
unsafe { ffi::flodl_nccl_destroy_rank(self.handle) };
self.handle = ptr::null_mut();
}
}
}
impl std::fmt::Debug for NcclRankComm {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("NcclRankComm")
.field("rank", &self.rank)
.field("world_size", &self.world_size)
.finish()
}
}
#[cfg(test)]
#[path = "nccl_tests.rs"]
mod tests;