//! Integration tests for weight loading. //! //! These tests require the generated weight files from convert_fno_weights.py. use rtx_backend_cpu::{CpuBackend, CpuDevice}; use rtx_neural_operator::{FNO2d, SafeTensorsFile, load_config, load_fno2d_weights}; use rtx_nn::generic::GenericModule4D; use rtx_tensor::generic::GenericTensor; use std::path::Path; const WEIGHTS_PATH: &str = "weights/fno/fno2d_darcy/model.safetensors"; const CONFIG_PATH: &str = "weights/fno/fno2d_darcy/config.json"; fn get_workspace_root() -> std::path::PathBuf { // Navigate from crate directory to workspace root let manifest_dir = std::env::var("CARGO_MANIFEST_DIR").unwrap(); Path::new(&manifest_dir) .parent() .unwrap() .parent() .unwrap() .parent() .unwrap() .to_path_buf() } #[test] fn test_load_safetensors_file() { let workspace = get_workspace_root(); let weights_path = workspace.join(WEIGHTS_PATH); if !weights_path.exists() { eprintln!( "Skipping test: weights file not found at {:?}", weights_path ); eprintln!( "Run: python scripts/convert_fno_weights.py --model fno2d_darcy --train --output weights/fno" ); return; } let safetensors = SafeTensorsFile::load(&weights_path).expect("Failed to load SafeTensors file"); // Check that expected tensors are present let tensor_names = safetensors.tensor_names(); println!("Found {} tensors", tensor_names.len()); for name in &tensor_names { println!(" {}", name); } // Check for either legacy or neuraloperator v2.0 format let has_legacy = safetensors.contains("lifting.weight"); let has_v2 = safetensors.contains("lifting.fcs.0.weight"); assert!( has_legacy || has_v2, "Missing lifting weights (neither legacy nor v2.0 format)" ); assert!( safetensors.contains("spectral_conv.0.weights1_real"), "Missing spectral_conv.0.weights1_real" ); assert!( safetensors.contains("projection.1.weight"), "Missing projection.1.weight" ); // Check lifting weight shape based on format if has_v2 { // neuraloperator v2.0 format: fc1 is [2*width, in_channels+2] let fc1 = safetensors.get("lifting.fcs.0.weight").unwrap(); let fc1_shape = fc1.shape(); println!("Lifting fc1 shape: {:?}", fc1_shape); // fc1 output should be 2*width = 64 assert_eq!(fc1_shape[0], 64, "Wrong fc1 output dimension"); // fc1 input should be in_channels + 2 = 3 (for darcy, data_channels=1 + 2 pos encoding) assert_eq!(fc1_shape[1], 3, "Wrong fc1 input dimension"); let fc2 = safetensors.get("lifting.fcs.1.weight").unwrap(); let fc2_shape = fc2.shape(); println!("Lifting fc2 shape: {:?}", fc2_shape); // fc2 should be [width, 2*width] assert_eq!(fc2_shape[0], 32, "Wrong fc2 output dimension"); assert_eq!(fc2_shape[1], 64, "Wrong fc2 input dimension"); } else { // Legacy format let lifting = safetensors.get("lifting.weight").unwrap(); assert_eq!(lifting.shape(), &[32, 3], "Wrong lifting weight shape"); } } #[test] fn test_load_config() { let workspace = get_workspace_root(); let config_path = workspace.join(CONFIG_PATH); if !config_path.exists() { eprintln!("Skipping test: config file not found at {:?}", config_path); return; } let config = load_config(&config_path).expect("Failed to load config"); assert_eq!(config.in_channels, 3); assert_eq!(config.out_channels, 1); assert_eq!(config.width, 32); assert_eq!(config.n_modes, (12, 12)); assert_eq!(config.n_layers, 4); assert_eq!(config.pde_type, "darcy"); } #[test] fn test_load_fno2d_weights() { let workspace = get_workspace_root(); let weights_path = workspace.join(WEIGHTS_PATH); if !weights_path.exists() { eprintln!( "Skipping test: weights file not found at {:?}", weights_path ); return; } // Check if weights are in neuraloperator v2.0 format let safetensors = SafeTensorsFile::load(&weights_path).expect("Failed to load SafeTensors"); let is_v2_format = safetensors.contains("lifting.fcs.0.weight"); if !is_v2_format { eprintln!("Skipping test: weights are in legacy format, not neuraloperator v2.0 format"); eprintln!("To generate v2.0 weights, run:"); eprintln!( " python3 scripts/convert_fno_weights.py --model fno2d_darcy --train --output weights/fno" ); return; } let weights = load_fno2d_weights(&weights_path).expect("Failed to load FNO2d weights"); // Check config was inferred correctly // Note: in_channels is the data channels (without positional encoding) // For Darcy, this is 1 (the coefficient field) println!( "Config: in_channels={}, out_channels={}, width={}, n_layers={}", weights.config.in_channels, weights.config.out_channels, weights.config.width, weights.config.n_layers ); assert_eq!(weights.config.out_channels, 1); assert_eq!(weights.config.width, 32); assert_eq!(weights.config.n_layers, 4); // Check lifting MLP weight shapes (neuraloperator v2.0 format) // fc1: [2*width, in_channels+2] = [64, in_channels+2] // fc2: [width, 2*width] = [32, 64] let hidden_dim = 2 * weights.config.width; // 64 let in_with_pos = weights.config.in_channels + 2; assert_eq!( weights.lifting_fc1_weight.len(), hidden_dim * in_with_pos, "fc1 weight wrong size: got {}, expected {}", weights.lifting_fc1_weight.len(), hidden_dim * in_with_pos ); assert_eq!( weights.lifting_fc1_bias.len(), hidden_dim, "fc1 bias wrong size" ); assert_eq!( weights.lifting_fc2_weight.len(), weights.config.width * hidden_dim, "fc2 weight wrong size: got {}, expected {}", weights.lifting_fc2_weight.len(), weights.config.width * hidden_dim ); assert_eq!( weights.lifting_fc2_bias.len(), weights.config.width, "fc2 bias wrong size" ); // Check other layers assert_eq!(weights.spectral_weights.len(), 4); assert_eq!(weights.conv_weights.len(), 4); assert_eq!(weights.projection_weights.len(), 2); // Check spectral conv weights let expected_spectral_size = 32 * 32 * 12 * 12; for (i, sw) in weights.spectral_weights.iter().enumerate() { assert_eq!( sw.weights1_real.len(), expected_spectral_size, "Layer {} weights1_real wrong size", i ); assert_eq!( sw.weights1_imag.len(), expected_spectral_size, "Layer {} weights1_imag wrong size", i ); } } #[test] fn test_fno2d_from_weights_and_inference() { let workspace = get_workspace_root(); let weights_path = workspace.join(WEIGHTS_PATH); if !weights_path.exists() { eprintln!( "Skipping test: weights file not found at {:?}", weights_path ); return; } // Check if weights are in neuraloperator v2.0 format let safetensors = SafeTensorsFile::load(&weights_path).expect("Failed to load SafeTensors"); let is_v2_format = safetensors.contains("lifting.fcs.0.weight"); if !is_v2_format { eprintln!("Skipping test: weights are in legacy format, not neuraloperator v2.0 format"); eprintln!("To generate v2.0 weights, run:"); eprintln!( " python3 scripts/convert_fno_weights.py --model fno2d_darcy --train --output weights/fno" ); return; } let weights = load_fno2d_weights(&weights_path).expect("Failed to load FNO2d weights"); let device = CpuDevice::new(); let fno = FNO2d::::from_weights(&weights, &device) .expect("Failed to create FNO2d from weights"); // Create input tensor using in_channels from config // The model expects [batch, in_channels, height, width] // Positional encoding is added internally, so we use in_channels (not in_channels+2) let in_channels = weights.config.in_channels; let input = GenericTensor::randn([1, in_channels, 64, 64], &device); // Run inference let output = fno.forward_4d(&input); // Check output shape assert_eq!(output.shape()[0], 1, "Wrong batch size"); assert_eq!( output.shape()[1], weights.config.out_channels, "Wrong output channels" ); assert_eq!(output.shape()[2], 64, "Wrong output height"); assert_eq!(output.shape()[3], 64, "Wrong output width"); // Check output is not all zeros let output_data = output.to_vec(); let non_zero = output_data.iter().filter(|&&x| x.abs() > 1e-10).count(); assert!(non_zero > 0, "Output is all zeros"); println!("FNO2d inference successful!"); println!(" Input shape: {:?}", input.shape()); println!(" Output shape: {:?}", output.shape()); println!( " Output range: [{:.4}, {:.4}]", output_data.iter().cloned().reduce(f32::min).unwrap(), output_data.iter().cloned().reduce(f32::max).unwrap() ); } const REFERENCE_INPUT_PATH: &str = "weights/fno/fno2d_darcy/reference_input.safetensors"; const REFERENCE_OUTPUT_PATH: &str = "weights/fno/fno2d_darcy/reference_output.safetensors"; use rtx_neural_operator::{LiftingMLP, SpectralConv2d}; /// Test that Rust FNO2d produces the same output as Python implementation. /// /// This test loads reference input/output generated by Python and verifies /// numerical equivalence within a tolerance. #[test] fn test_python_rust_numerical_equivalence() { let workspace = get_workspace_root(); let weights_path = workspace.join(WEIGHTS_PATH); let input_path = workspace.join(REFERENCE_INPUT_PATH); let output_path = workspace.join(REFERENCE_OUTPUT_PATH); // Check all required files exist if !weights_path.exists() { eprintln!( "Skipping test: weights file not found at {:?}", weights_path ); eprintln!("Run: python scripts/generate_reference_outputs.py"); return; } if !input_path.exists() { eprintln!( "Skipping test: reference input not found at {:?}", input_path ); eprintln!("Run: python scripts/generate_reference_outputs.py"); return; } if !output_path.exists() { eprintln!( "Skipping test: reference output not found at {:?}", output_path ); eprintln!("Run: python scripts/generate_reference_outputs.py"); return; } // Check v2.0 format let safetensors = SafeTensorsFile::load(&weights_path).expect("Failed to load weights"); if !safetensors.contains("lifting.fcs.0.weight") { eprintln!("Skipping test: weights are in legacy format"); return; } // Load weights and create model let weights = load_fno2d_weights(&weights_path).expect("Failed to load FNO2d weights"); let device = CpuDevice::new(); let fno = FNO2d::::from_weights(&weights, &device) .expect("Failed to create FNO2d from weights"); // Load reference input let input_safetensors = SafeTensorsFile::load(&input_path).expect("Failed to load reference input"); let input_tensor_info = input_safetensors .get("input") .expect("Reference input missing 'input' tensor"); let input_data = input_tensor_info .to_f32() .expect("Failed to convert input to f32"); let input_shape = input_tensor_info.shape(); println!("Reference input shape: {:?}", input_shape); // Create input tensor from reference data let input = GenericTensor::from_slice( &input_data, [ input_shape[0], input_shape[1], input_shape[2], input_shape[3], ], &device, ); // Run Rust inference let rust_output = fno.forward_4d(&input); let rust_output_data = rust_output.to_vec(); // Load reference output let output_safetensors = SafeTensorsFile::load(&output_path).expect("Failed to load reference output"); let output_tensor_info = output_safetensors .get("output") .expect("Reference output missing 'output' tensor"); let python_output_data = output_tensor_info .to_f32() .expect("Failed to convert output to f32"); let python_output_shape = output_tensor_info.shape(); println!("Reference output shape: {:?}", python_output_shape); println!("Rust output shape: {:?}", rust_output.shape()); // Check shapes match assert_eq!( rust_output.shape()[0], python_output_shape[0], "Batch size mismatch" ); assert_eq!( rust_output.shape()[1], python_output_shape[1], "Channel count mismatch" ); assert_eq!( rust_output.shape()[2], python_output_shape[2], "Height mismatch" ); assert_eq!( rust_output.shape()[3], python_output_shape[3], "Width mismatch" ); // Compare outputs let n_elements = rust_output_data.len(); assert_eq!(n_elements, python_output_data.len(), "Output size mismatch"); // Calculate statistics let mut max_abs_diff: f32 = 0.0; let mut sum_abs_diff: f32 = 0.0; let mut sum_sq_diff: f32 = 0.0; let mut max_rel_diff: f32 = 0.0; for (i, (rust_val, python_val)) in rust_output_data .iter() .zip(python_output_data.iter()) .enumerate() { let abs_diff = (rust_val - python_val).abs(); sum_abs_diff += abs_diff; sum_sq_diff += abs_diff * abs_diff; if abs_diff > max_abs_diff { max_abs_diff = abs_diff; } // Relative difference (avoid division by zero) let rel_diff = if python_val.abs() > 1e-10 { abs_diff / python_val.abs() } else { abs_diff }; if rel_diff > max_rel_diff { max_rel_diff = rel_diff; } // Print first few differences for debugging if i < 5 { println!(" [{i}] rust={rust_val:.6}, python={python_val:.6}, diff={abs_diff:.6e}"); } } let mean_abs_diff = sum_abs_diff / n_elements as f32; let rmse = (sum_sq_diff / n_elements as f32).sqrt(); println!("\nNumerical comparison:"); println!(" Max absolute difference: {:.6e}", max_abs_diff); println!(" Mean absolute difference: {:.6e}", mean_abs_diff); println!(" RMSE: {:.6e}", rmse); println!(" Max relative difference: {:.6e}", max_rel_diff); // Python output statistics let python_mean: f32 = python_output_data.iter().sum::() / n_elements as f32; let python_min = python_output_data.iter().cloned().reduce(f32::min).unwrap(); let python_max = python_output_data.iter().cloned().reduce(f32::max).unwrap(); println!( "\nPython output: mean={:.6}, range=[{:.6}, {:.6}]", python_mean, python_min, python_max ); // Rust output statistics let rust_mean: f32 = rust_output_data.iter().sum::() / n_elements as f32; let rust_min = rust_output_data.iter().cloned().reduce(f32::min).unwrap(); let rust_max = rust_output_data.iter().cloned().reduce(f32::max).unwrap(); println!( "Rust output: mean={:.6}, range=[{:.6}, {:.6}]", rust_mean, rust_min, rust_max ); // Assert numerical equivalence within tolerance // FFT operations can accumulate small numerical differences, so we use relaxed tolerance let tolerance = 1e-4; assert!( max_abs_diff < tolerance, "Max absolute difference {:.6e} exceeds tolerance {:.6e}", max_abs_diff, tolerance ); println!("\nPython vs Rust comparison PASSED!"); println!( "Maximum difference {:.6e} is within tolerance {:.6e}", max_abs_diff, tolerance ); } /// Test isolated spectral convolution layer against Python reference. #[test] fn test_spectral_conv_isolated() { // Load debug data from Python script let debug_path = std::path::Path::new("/tmp/spectral_conv_debug.safetensors"); if !debug_path.exists() { eprintln!("Skipping test: debug data not found at {:?}", debug_path); eprintln!("Run: python scripts/debug_spectral_conv.py"); return; } let safetensors = SafeTensorsFile::load(debug_path).expect("Failed to load debug data"); // Load tensors let input_info = safetensors.get("input").expect("Missing input"); let w1_real_info = safetensors.get("w1_real").expect("Missing w1_real"); let w1_imag_info = safetensors.get("w1_imag").expect("Missing w1_imag"); let w2_real_info = safetensors.get("w2_real").expect("Missing w2_real"); let w2_imag_info = safetensors.get("w2_imag").expect("Missing w2_imag"); let expected_info = safetensors.get("output").expect("Missing output"); let input_data = input_info.to_f32().expect("Failed to convert input"); let w1_real = w1_real_info.to_f32().expect("Failed to convert w1_real"); let w1_imag = w1_imag_info.to_f32().expect("Failed to convert w1_imag"); let w2_real = w2_real_info.to_f32().expect("Failed to convert w2_real"); let w2_imag = w2_imag_info.to_f32().expect("Failed to convert w2_imag"); let expected_data = expected_info.to_f32().expect("Failed to convert output"); let input_shape = input_info.shape(); let w1_shape = w1_real_info.shape(); println!("Input shape: {:?}", input_shape); println!("W1 shape: {:?}", w1_shape); println!("Expected output shape: {:?}", expected_info.shape()); let batch = input_shape[0]; let in_ch = input_shape[1]; let height = input_shape[2]; let width = input_shape[3]; let out_ch = w1_shape[1]; let modes_h = w1_shape[2]; let modes_w = w1_shape[3]; println!( "batch={}, in_ch={}, out_ch={}, height={}, width={}", batch, in_ch, out_ch, height, width ); println!("modes_h={}, modes_w={}", modes_h, modes_w); // Create spectral conv layer with the debug weights let device = CpuDevice::new(); let spectral_conv = SpectralConv2d::::from_weights_dual( &w1_real, &w1_imag, &w2_real, &w2_imag, in_ch, out_ch, modes_h, modes_w, &device, ); // Create input tensor let input = GenericTensor::from_slice(&input_data, [batch, in_ch, height, width], &device); // Run spectral conv let output = spectral_conv.forward_4d(&input); let output_data = output.to_vec(); println!( "\nPython output: range=[{:.6}, {:.6}], first 4: {:?}", expected_data.iter().cloned().reduce(f32::min).unwrap(), expected_data.iter().cloned().reduce(f32::max).unwrap(), &expected_data[..4] ); println!( "Rust output: range=[{:.6}, {:.6}], first 4: {:?}", output_data.iter().cloned().reduce(f32::min).unwrap(), output_data.iter().cloned().reduce(f32::max).unwrap(), &output_data[..4] ); // Calculate difference let mut max_diff: f32 = 0.0; for (r, p) in output_data.iter().zip(expected_data.iter()) { let diff = (r - p).abs(); if diff > max_diff { max_diff = diff; } } println!("\nMax absolute difference: {:.6e}", max_diff); // Check if outputs have similar magnitude (within 10x) let python_max = expected_data .iter() .map(|x| x.abs()) .reduce(f32::max) .unwrap(); let rust_max = output_data .iter() .map(|x| x.abs()) .reduce(f32::max) .unwrap(); let ratio = rust_max / python_max; println!("Magnitude ratio (Rust/Python): {:.4}", ratio); } /// Test isolated lifting MLP against Python reference. #[test] fn test_lifting_mlp_isolated() { // Load debug data from Python script let debug_path = std::path::Path::new("/tmp/lifting_mlp_debug.safetensors"); if !debug_path.exists() { eprintln!("Skipping test: debug data not found at {:?}", debug_path); eprintln!("Run: python scripts/debug_lifting_mlp.py"); return; } let safetensors = SafeTensorsFile::load(debug_path).expect("Failed to load debug data"); // Load tensors let input_info = safetensors.get("input").expect("Missing input"); let fc1_weight_info = safetensors.get("fc1_weight").expect("Missing fc1_weight"); let fc1_bias_info = safetensors.get("fc1_bias").expect("Missing fc1_bias"); let fc2_weight_info = safetensors.get("fc2_weight").expect("Missing fc2_weight"); let fc2_bias_info = safetensors.get("fc2_bias").expect("Missing fc2_bias"); let expected_info = safetensors.get("output").expect("Missing output"); let input_data = input_info.to_f32().expect("Failed to convert input"); let fc1_weight = fc1_weight_info .to_f32() .expect("Failed to convert fc1_weight"); let fc1_bias = fc1_bias_info.to_f32().expect("Failed to convert fc1_bias"); let fc2_weight = fc2_weight_info .to_f32() .expect("Failed to convert fc2_weight"); let fc2_bias = fc2_bias_info.to_f32().expect("Failed to convert fc2_bias"); let expected_data = expected_info.to_f32().expect("Failed to convert output"); let input_shape = input_info.shape(); let fc1_shape = fc1_weight_info.shape(); let fc2_shape = fc2_weight_info.shape(); println!("Input shape: {:?}", input_shape); println!("FC1 weight shape: {:?}", fc1_shape); println!("FC2 weight shape: {:?}", fc2_shape); println!("Expected output shape: {:?}", expected_info.shape()); let batch = input_shape[0]; let in_channels = input_shape[1]; let height = input_shape[2]; let width = input_shape[3]; let hidden_dim = fc1_shape[0]; let out_channels = fc2_shape[0]; println!( "in_channels={}, hidden_dim={}, out_channels={}", in_channels, hidden_dim, out_channels ); // Create lifting MLP with the debug weights let device = CpuDevice::new(); let lifting = LiftingMLP::::from_weights( &fc1_weight, &fc1_bias, &fc2_weight, &fc2_bias, in_channels, out_channels, &device, ); // Create input tensor let input = GenericTensor::from_slice(&input_data, [batch, in_channels, height, width], &device); // Run lifting MLP let output = lifting.forward_4d(&input); let output_data = output.to_vec(); println!( "\nPython output: range=[{:.4}, {:.4}], first 4: {:?}", expected_data.iter().cloned().reduce(f32::min).unwrap(), expected_data.iter().cloned().reduce(f32::max).unwrap(), &expected_data[..4] ); println!( "Rust output: range=[{:.4}, {:.4}], first 4: {:?}", output_data.iter().cloned().reduce(f32::min).unwrap(), output_data.iter().cloned().reduce(f32::max).unwrap(), &output_data[..4] ); // Calculate difference let mut max_diff: f32 = 0.0; for (r, p) in output_data.iter().zip(expected_data.iter()) { let diff: f32 = (r - p).abs(); if diff > max_diff { max_diff = diff; } } println!("\nMax absolute difference: {:.6e}", max_diff); // Check magnitude ratio let python_max: f32 = expected_data .iter() .map(|x| x.abs()) .reduce(f32::max) .unwrap(); let rust_max: f32 = output_data .iter() .map(|x| x.abs()) .reduce(f32::max) .unwrap(); let ratio = rust_max / python_max; println!("Magnitude ratio (Rust/Python): {:.4}", ratio); // GELU approximation can have small differences assert!( max_diff < 1e-3, "Max difference {} exceeds tolerance", max_diff ); }