Files
rustytorch/crates/training/rtx-flash-attention/src/error.rs
T
2026-03-04 00:08:42 +00:00

277 lines
9.2 KiB
Rust

//! Error types for Flash Attention operations
use thiserror::Error;
use rtx_tensor::TensorError;
/// Result type for Flash Attention operations
pub type FlashResult<T> = Result<T, FlashError>;
/// Flash Attention error types
#[derive(Error, Debug)]
pub enum FlashError {
/// CUDA runtime errors
#[error("CUDA error: {message}")]
Cuda { message: String },
/// Memory allocation errors
#[error("Memory allocation error: {message}")]
Memory { message: String },
/// Invalid configuration errors
#[error("Invalid configuration: {message}")]
Config { message: String },
/// Tensor operation errors
#[error("Tensor error: {message}")]
Tensor { message: String },
/// Tensor shape mismatch errors
#[error("Tensor shape mismatch: expected {expected:?}, got {actual:?}")]
ShapeMismatch {
expected: Vec<usize>,
actual: Vec<usize>,
},
/// Unsupported operation errors
#[error("Unsupported operation: {operation}")]
Unsupported { operation: String },
/// Kernel compilation errors
#[error("Kernel compilation error: {message}")]
KernelCompilation { message: String },
/// Backend initialization errors
#[error("Backend initialization error: {message}")]
BackendInit { message: String },
/// Edge deployment errors
#[cfg(feature = "edge")]
#[error("Edge deployment error: {message}")]
Edge { message: String },
/// Numerical stability errors
#[error("Numerical instability detected: {message}")]
NumericalInstability { message: String },
/// Hardware compatibility errors
#[error("Hardware compatibility error: {message}")]
HardwareCompat { message: String },
/// Generic I/O errors
#[error("I/O error: {0}")]
Io(#[from] std::io::Error),
/// Generic errors from other crates
#[error("External error: {0}")]
External(#[from] anyhow::Error),
}
impl FlashError {
/// Create a CUDA error
pub fn cuda<S: Into<String>>(message: S) -> Self {
Self::Cuda {
message: message.into(),
}
}
/// Create a memory error
pub fn memory<S: Into<String>>(message: S) -> Self {
Self::Memory {
message: message.into(),
}
}
/// Create a configuration error
pub fn config<S: Into<String>>(message: S) -> Self {
Self::Config {
message: message.into(),
}
}
/// Create a tensor error
pub fn tensor<S: Into<String>>(message: S) -> Self {
Self::Tensor {
message: message.into(),
}
}
/// Create a shape mismatch error
pub fn shape_mismatch(expected: Vec<usize>, actual: Vec<usize>) -> Self {
Self::ShapeMismatch { expected, actual }
}
/// Create an unsupported operation error
pub fn unsupported<S: Into<String>>(operation: S) -> Self {
Self::Unsupported {
operation: operation.into(),
}
}
/// Create a kernel compilation error
pub fn kernel_compilation<S: Into<String>>(message: S) -> Self {
Self::KernelCompilation {
message: message.into(),
}
}
/// Create a backend initialization error
pub fn backend_init<S: Into<String>>(message: S) -> Self {
Self::BackendInit {
message: message.into(),
}
}
/// Create a numerical instability error
pub fn numerical_instability<S: Into<String>>(message: S) -> Self {
Self::NumericalInstability {
message: message.into(),
}
}
/// Create a hardware compatibility error
pub fn hardware_compat<S: Into<String>>(message: S) -> Self {
Self::HardwareCompat {
message: message.into(),
}
}
#[cfg(feature = "edge")]
/// Create an edge deployment error
pub fn edge<S: Into<String>>(message: S) -> Self {
Self::Edge {
message: message.into(),
}
}
/// Check if this is a recoverable error
pub fn is_recoverable(&self) -> bool {
match self {
Self::Cuda { .. } => false, // CUDA errors usually require restart
Self::Memory { .. } => true, // Can try with smaller batches
Self::Config { .. } => false, // Configuration errors need fixing
Self::ShapeMismatch { .. } => false, // Shape errors need fixing
Self::Unsupported { .. } => false, // Unsupported operations won't work
Self::KernelCompilation { .. } => false, // Compilation errors need fixing
Self::BackendInit { .. } => false, // Backend errors need restart
Self::NumericalInstability { .. } => true, // Can retry with different params
Self::HardwareCompat { .. } => false, // Hardware issues won't resolve
Self::Io(_) => true, // I/O errors might be transient
Self::External(_) => true, // External errors might be recoverable
Self::Tensor { .. } => true, // Tensor errors might be recoverable
#[cfg(feature = "edge")]
FlashError::Edge { .. } => true, // Edge errors might be transient
}
}
/// Get error category for logging/metrics
pub fn category(&self) -> &'static str {
match self {
Self::Cuda { .. } => "cuda",
Self::Memory { .. } => "memory",
Self::Config { .. } => "config",
Self::ShapeMismatch { .. } => "shape",
Self::Unsupported { .. } => "unsupported",
Self::KernelCompilation { .. } => "kernel",
Self::BackendInit { .. } => "backend",
Self::NumericalInstability { .. } => "numerical",
Self::HardwareCompat { .. } => "hardware",
Self::Io(_) => "io",
Self::External(_) => "external",
Self::Tensor { .. } => "tensor",
#[cfg(feature = "edge")]
FlashError::Edge { .. } => "edge",
}
}
}
// Convert from TensorError to FlashError
impl From<TensorError> for FlashError {
fn from(err: TensorError) -> Self {
Self::Tensor {
message: err.to_string(),
}
}
}
// Convert from RuntimeError to FlashError
impl From<rtx_runtime::RuntimeError> for FlashError {
fn from(err: rtx_runtime::RuntimeError) -> Self {
use rtx_runtime::RuntimeError;
match err {
RuntimeError::InvalidOperation(msg) => {
Self::Cuda { message: format!("Runtime operation error: {msg}") }
},
RuntimeError::AllocationFailed { message, size, available } => {
Self::Memory {
message: format!(
"Memory allocation failed: {message} (requested: {size} bytes, available: {available} bytes)"
)
}
},
RuntimeError::FragmentationExceeded { threshold, current } => {
Self::Memory {
message: format!(
"Memory fragmentation exceeded {threshold}%: current fragmentation {current}%"
)
}
},
RuntimeError::DeviceError { device_id, message } => {
Self::Cuda { message: format!("Device {device_id} error: {message}") }
},
RuntimeError::StreamError { stream_id, message } => {
Self::Cuda { message: format!("Stream {stream_id} error: {message}") }
},
RuntimeError::KernelError { message } => {
Self::Cuda { message: format!("Kernel execution error: {message}") }
},
RuntimeError::ConfigError { message } => {
Self::Config { message: format!("Runtime config error: {message}") }
},
RuntimeError::ResourceExhausted { resource } => {
Self::Memory { message: format!("Resource exhausted: {resource}") }
},
RuntimeError::BackendError { backend, source } => {
Self::Cuda { message: format!("Backend error ({backend}): {source}") }
},
RuntimeError::Fusion { message } => {
Self::KernelCompilation { message: format!("Kernel fusion error: {message}") }
},
RuntimeError::OptimizationError(msg) => {
Self::Cuda { message: format!("Optimization error: {msg}") }
},
RuntimeError::BackendNotSupported(backend) => {
Self::Unsupported { operation: format!("Backend not supported: {backend}") }
},
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_error_creation() {
let err = FlashError::cuda("test message");
assert!(matches!(err, FlashError::Cuda { .. }));
assert_eq!(err.category(), "cuda");
assert!(!err.is_recoverable());
}
#[test]
fn test_shape_mismatch() {
let err = FlashError::shape_mismatch(vec![2, 3, 4], vec![2, 3, 5]);
assert!(matches!(err, FlashError::ShapeMismatch { .. }));
assert_eq!(err.category(), "shape");
}
#[test]
fn test_recoverable_errors() {
assert!(FlashError::memory("test").is_recoverable());
assert!(!FlashError::config("test").is_recoverable());
assert!(FlashError::numerical_instability("test").is_recoverable());
}
}