Initial commit
This commit is contained in:
@@ -0,0 +1,98 @@
|
||||
//! GPU kernels for NMF operations
|
||||
//!
|
||||
//! This module provides GPU implementations for Non-negative Matrix
|
||||
//! Factorization operations, supporting both CUDA (NVIDIA) and
|
||||
//! Metal (Apple Silicon) backends.
|
||||
//!
|
||||
//! ## Operations
|
||||
//!
|
||||
//! - **Multiplicative updates**: Fused operations for W and H matrix updates
|
||||
//! - **Non-negativity clamping**: Ensure all values >= 0
|
||||
//! - **Frobenius norm**: Convergence checking via ||A||_F calculation
|
||||
//! - **Reconstruction error**: Compute ||V - WH||_F^2
|
||||
//!
|
||||
//! ## Backend Selection
|
||||
//!
|
||||
//! - On macOS with Metal feature: Uses Metal compute shaders
|
||||
//! - On Linux/Windows with CUDA feature: Uses CUDA kernels
|
||||
//! - Otherwise: Falls back to CPU implementation
|
||||
|
||||
// CUDA kernels (Linux/Windows with NVIDIA GPU)
|
||||
#[cfg(feature = "cuda")]
|
||||
pub mod custom_kernels;
|
||||
|
||||
#[cfg(feature = "cuda")]
|
||||
pub use custom_kernels::NMFKernels as CudaNMFKernels;
|
||||
|
||||
// Metal kernels (macOS with Apple Silicon)
|
||||
#[cfg(all(target_os = "macos", feature = "metal"))]
|
||||
pub mod metal_kernels;
|
||||
|
||||
#[cfg(all(target_os = "macos", feature = "metal"))]
|
||||
pub use metal_kernels::{ElementwiseParams, MetalNMFKernels, NMFUpdateParams, ReductionParams};
|
||||
|
||||
/// Check if GPU acceleration is available
|
||||
pub fn is_gpu_available() -> bool {
|
||||
#[cfg(all(target_os = "macos", feature = "metal"))]
|
||||
{
|
||||
return metal_kernels::MetalNMFKernels::is_available();
|
||||
}
|
||||
|
||||
#[cfg(feature = "cuda")]
|
||||
{
|
||||
// Check CUDA availability (simplified)
|
||||
return true;
|
||||
}
|
||||
|
||||
#[allow(unreachable_code)]
|
||||
false
|
||||
}
|
||||
|
||||
/// GPU backend type
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub enum GpuBackend {
|
||||
/// Apple Metal (macOS)
|
||||
Metal,
|
||||
/// NVIDIA CUDA
|
||||
Cuda,
|
||||
/// CPU fallback
|
||||
Cpu,
|
||||
}
|
||||
|
||||
/// Get the current GPU backend
|
||||
pub fn current_backend() -> GpuBackend {
|
||||
#[cfg(all(target_os = "macos", feature = "metal"))]
|
||||
{
|
||||
if metal_kernels::MetalNMFKernels::is_available() {
|
||||
return GpuBackend::Metal;
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(feature = "cuda")]
|
||||
{
|
||||
return GpuBackend::Cuda;
|
||||
}
|
||||
|
||||
#[allow(unreachable_code)]
|
||||
GpuBackend::Cpu
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn test_kernels_module_exists() {
|
||||
// Basic test to ensure module compiles
|
||||
assert!(true);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_backend_detection() {
|
||||
let backend = current_backend();
|
||||
println!("Current NMF GPU backend: {:?}", backend);
|
||||
|
||||
let gpu_available = is_gpu_available();
|
||||
println!("GPU available: {}", gpu_available);
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user