//! Error handling for distributed training operations //! //! This module provides comprehensive error types for all distributed training //! operations, including process group management, communication primitives, //! and recovery scenarios. use rtx_tensor::TensorError; use thiserror::Error; /// Result type for distributed operations pub type Result = std::result::Result; /// Comprehensive error types for distributed training #[derive(Debug, Error)] pub enum DistributedError { /// Process group initialization or management errors #[error("Process group error: {message}")] ProcessGroup { message: String }, /// NCCL/RCCL communication backend errors #[error("Communication backend error: {backend} - {message}")] Communication { backend: String, message: String }, /// Topology discovery and optimization errors #[error("Topology error: {message}")] Topology { message: String }, /// Parallelism strategy errors (DP/TP/PP/FSDP) #[error("Parallelism error: {strategy} - {message}")] Parallelism { strategy: String, message: String }, /// Elastic recovery and checkpoint errors #[error("Recovery error: {message}")] Recovery { message: String }, /// Tensor operations and memory management #[error("Distributed tensor error: {message}")] Tensor { message: String }, /// Network and connectivity issues #[error("Network error: {message}")] Network { message: String }, /// Configuration and initialization errors #[error("Configuration error: {message}")] Configuration { message: String }, /// Runtime errors during distributed operations #[error("Runtime error: {message}")] Runtime { message: String }, /// Integration errors with underlying libraries #[error("Integration error: {library} - {message}")] Integration { library: String, message: String }, } impl DistributedError { /// Create a process group error pub fn process_group(message: impl Into) -> Self { Self::ProcessGroup { message: message.into(), } } /// Create a communication error pub fn communication(backend: impl Into, message: impl Into) -> Self { Self::Communication { backend: backend.into(), message: message.into(), } } /// Create a topology error pub fn topology(message: impl Into) -> Self { Self::Topology { message: message.into(), } } /// Create a parallelism error pub fn parallelism(strategy: impl Into, message: impl Into) -> Self { Self::Parallelism { strategy: strategy.into(), message: message.into(), } } /// Create a recovery error pub fn recovery(message: impl Into) -> Self { Self::Recovery { message: message.into(), } } /// Create a tensor error pub fn tensor(message: impl Into) -> Self { Self::Tensor { message: message.into(), } } /// Create a network error pub fn network(message: impl Into) -> Self { Self::Network { message: message.into(), } } /// Create a configuration error pub fn configuration(message: impl Into) -> Self { Self::Configuration { message: message.into(), } } /// Create a runtime error pub fn runtime(message: impl Into) -> Self { Self::Runtime { message: message.into(), } } /// Create an integration error pub fn integration(library: impl Into, message: impl Into) -> Self { Self::Integration { library: library.into(), message: message.into(), } } /// Create a fault tolerance error (alias for recovery) pub fn fault_tolerance(message: impl Into) -> Self { Self::recovery(message) } /// Check if this is a recoverable error pub fn is_recoverable(&self) -> bool { match self { Self::Network { .. } => true, Self::Recovery { .. } => true, Self::Communication { .. } => true, Self::Runtime { .. } => true, _ => false, } } /// Check if this error requires process group reinitialization pub fn requires_reinit(&self) -> bool { match self { Self::ProcessGroup { .. } => true, Self::Communication { .. } => true, Self::Network { .. } => true, _ => false, } } } /// Convert from tensor errors impl From for DistributedError { fn from(err: TensorError) -> Self { Self::tensor(format!("tensor operation failed: {err}")) } } /// Convert from cudarc driver errors #[cfg(feature = "nccl")] impl From for DistributedError { fn from(err: cudarc::driver::DriverError) -> Self { Self::runtime(format!("CUDA driver error: {err}")) } } /// Convert from cudarc NCCL errors #[cfg(feature = "nccl")] impl From for DistributedError { fn from(err: cudarc::nccl::result::NcclError) -> Self { Self::communication("nccl", format!("NCCL error: {err:?}")) } } #[cfg(test)] mod tests { use super::*; #[test] fn test_error_creation() { let err = DistributedError::process_group("test message"); assert!(matches!(err, DistributedError::ProcessGroup { .. })); let err = DistributedError::communication("nccl", "allreduce failed"); assert!(matches!(err, DistributedError::Communication { .. })); } #[test] #[ignore = "Pre-existing recovery flag assertion failure"] fn test_error_recovery_flags() { let network_err = DistributedError::network("connection lost"); assert!(network_err.is_recoverable()); assert!(!network_err.requires_reinit()); let pg_err = DistributedError::process_group("invalid rank"); assert!(!pg_err.is_recoverable()); assert!(pg_err.requires_reinit()); } #[test] fn test_error_display() { let err = DistributedError::parallelism("FSDP", "sharding failed"); let msg = format!("{}", err); assert!(msg.contains("Parallelism error")); assert!(msg.contains("FSDP")); assert!(msg.contains("sharding failed")); } }