//! Simple CUDA functionality test to verify compatibility //! //! This test verifies that the updated candle-core with CUDA support //! can successfully create tensors and perform basic operations. #![allow(unused_crate_dependencies)] use candle_core::{Device, Tensor}; use candle_nn::{linear, Module, VarBuilder, VarMap}; /// Test basic CUDA tensor operations pub fn test_cuda_basic() -> Result<(), Box> { println!("Testing CUDA compatibility..."); // Try to get CUDA device match Device::new_cuda(0) { Ok(device) => { println!("✅ CUDA device 0 available"); // Create a simple tensor let tensor = Tensor::randn(0f32, 1.0, (4, 4), &device)?; println!("✅ Created CUDA tensor: {:?}", tensor.shape()); // Perform basic operations let result = tensor.matmul(&tensor.t()?)?; println!("✅ Matrix multiplication successful: {:?}", result.shape()); Ok(()) }, Err(e) => { println!("⚠️ CUDA device not available: {}", e); println!("This is expected if no GPU is present, but CUDA compilation succeeded"); Ok(()) }, } } /// Test candle-nn components with CUDA pub fn test_cuda_neural_network() -> Result<(), Box> { println!("Testing CUDA neural network components..."); match Device::new_cuda(0) { Ok(device) => { let varmap = VarMap::new(); let vs = VarBuilder::from_varmap(&varmap, candle_core::DType::F32, &device); // Create a simple linear layer let linear_layer = linear(10, 5, vs.pp("linear"))?; let input = Tensor::randn(0f32, 1.0, (1, 10), &device)?; let output = linear_layer.forward(&input)?; println!( "✅ Neural network forward pass successful: {:?}", output.shape() ); Ok(()) }, Err(_) => { println!("⚠️ CUDA neural network test skipped (no GPU)"); Ok(()) }, } } #[cfg(test)] mod tests { use super::*; #[test] fn verify_cuda_compilation() { // This test just verifies that CUDA code compiles // The actual runtime test is optional since CI may not have GPU println!("CUDA compilation test passed!"); // Try to run basic test but don't fail if no GPU let _ = test_cuda_basic(); let _ = test_cuda_neural_network(); } } fn main() -> Result<(), Box> { test_cuda_basic()?; test_cuda_neural_network()?; println!("🎉 CUDA compatibility verification complete!"); Ok(()) }