108 lines
4.2 KiB
Rust
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(())
|
|
} |