Files
rustytorch/tests/validation/validate_ifft_functionality.rs
T
2026-03-04 00:08:42 +00:00

108 lines
4.2 KiB
Rust

//! Validation script for IFFT functionality
//!
//! This script validates that IFFT operations are working correctly in the RTX ecosystem.
use rtx_tensor::{ComplexTensor, Device, Tensor};
fn main() -> rtx_tensor::Result<()> {
println!("🔄 Validating IFFT Functionality...");
let device = Device::cpu();
// Test 1: Basic IFFT functionality
println!("\n📊 Test 1: Basic IFFT");
// Create a simple constant signal in frequency domain
let real_tensor = Tensor::from_data(vec![4.0, 0.0, 0.0, 0.0], vec![4], &device)?;
let imag_tensor = Tensor::from_data(vec![0.0, 0.0, 0.0, 0.0], vec![4], &device)?;
let fft_result = ComplexTensor::<f32>::from_real_imag(real_tensor, imag_tensor)?;
let ifft_result = fft_result.ifft()?;
let real_data = ifft_result.real().to_cpu()?;
let imag_data = ifft_result.imag().to_cpu()?;
println!(" Input FFT (real): {:?}", fft_result.real().to_cpu()?);
println!(" IFFT result (real): {:?}", real_data);
println!(" IFFT result (imag): {:?}", imag_data);
// Verify the result - should be constant 1.0
let expected = vec![1.0, 1.0, 1.0, 1.0];
let tolerance = 1e-5;
for (i, (&actual, &expected_val)) in real_data.iter().zip(&expected).enumerate() {
if (actual - expected_val).abs() > tolerance {
println!(" ❌ IFFT Test Failed: element {} = {}, expected {}", i, actual, expected_val);
return Ok(());
}
}
println!(" ✅ Basic IFFT test passed!");
// Test 2: Round-trip FFT->IFFT
println!("\n📊 Test 2: FFT->IFFT Round-trip");
let original_real = vec![1.0, 2.0, 3.0, 4.0];
let original_imag = vec![0.5, -0.5, 1.0, -1.0];
let real_tensor = Tensor::from_data(original_real.clone(), vec![4], &device)?;
let imag_tensor = Tensor::from_data(original_imag.clone(), vec![4], &device)?;
let original = ComplexTensor::<f32>::from_real_imag(real_tensor, imag_tensor)?;
// Forward FFT then inverse FFT
let fft_result = original.fft()?;
let recovered = fft_result.ifft()?;
let recovered_real = recovered.real().to_cpu()?;
let recovered_imag = recovered.imag().to_cpu()?;
println!(" Original (real): {:?}", original_real);
println!(" Recovered (real): {:?}", recovered_real);
println!(" Original (imag): {:?}", original_imag);
println!(" Recovered (imag): {:?}", recovered_imag);
// Verify round-trip accuracy
for i in 0..4 {
if (original_real[i] - recovered_real[i]).abs() > tolerance {
println!(" ❌ Round-trip failed (real): element {} = {}, expected {}",
i, recovered_real[i], original_real[i]);
return Ok(());
}
if (original_imag[i] - recovered_imag[i]).abs() > tolerance {
println!(" ❌ Round-trip failed (imag): element {} = {}, expected {}",
i, recovered_imag[i], original_imag[i]);
return Ok(());
}
}
println!(" ✅ Round-trip test passed!");
// Test 3: 2D IFFT
println!("\n📊 Test 3: 2D IFFT");
let real_tensor = Tensor::from_data(vec![1.0, 2.0, 3.0, 4.0], vec![2, 2], &device)?;
let imag_tensor = Tensor::from_data(vec![0.0, 0.0, 0.0, 0.0], vec![2, 2], &device)?;
let tensor_2d = ComplexTensor::<f32>::from_real_imag(real_tensor, imag_tensor)?;
// Test 2D FFT->IFFT round-trip
let fft_2d = tensor_2d.fft2d()?;
let recovered_2d = fft_2d.ifft2()?;
let original_data = tensor_2d.real().to_cpu()?;
let recovered_data = recovered_2d.real().to_cpu()?;
println!(" Original 2D (real): {:?}", original_data);
println!(" Recovered 2D (real): {:?}", recovered_data);
for i in 0..4 {
if (original_data[i] - recovered_data[i]).abs() > tolerance {
println!(" ❌ 2D IFFT failed: element {} = {}, expected {}",
i, recovered_data[i], original_data[i]);
return Ok(());
}
}
println!(" ✅ 2D IFFT test passed!");
println!("\n🎉 All IFFT tests passed! ✅");
println!(" - Basic 1D IFFT working");
println!(" - FFT->IFFT round-trip accurate");
println!(" - 2D IFFT operational");
Ok(())
}