1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122
| """ radarODE-MTL 测试脚本
测试内容: 1. 模型前向传播 2. ECG重建质量评估 3. R峰检测准确率 4. 噪声鲁棒性测试 """
import torch import numpy as np from scipy.signal import find_peaks from scipy.stats import pearsonr
def test_forward_pass(): """测试前向传播""" model = RadarODEMTL(radar_channels=8, ecg_length=256) model.eval() radar = torch.randn(4, 8, 1024) ecg, anchor, cycle = model(radar) assert ecg.shape == (4, 256), f"ECG shape error: {ecg.shape}" assert anchor.shape == (4, 1024), f"Anchor shape error: {anchor.shape}" assert cycle.shape == (4, 1024), f"Cycle shape error: {cycle.shape}" print("✓ Forward pass test passed") return model
def test_ecg_quality(): """测试ECG重建质量""" model = test_forward_pass() t = np.linspace(0, 1, 256) ecg_gt = ( 0.1 * np.exp(-(t - 0.1)**2 / 0.001) + -0.3 * np.exp(-(t - 0.2)**2 / 0.0005) + 1.0 * np.exp(-(t - 0.25)**2 / 0.0005) + -0.4 * np.exp(-(t - 0.3)**2 / 0.0005) + 0.3 * np.exp(-(t - 0.4)**2 / 0.002) + ) radar_sim = np.random.randn(8, 1024) * 0.1 for i in range(8): radar_sim[i] += np.interp(np.linspace(0, 1, 1024), t, ecg_gt) * (0.5 + 0.1*i) radar_tensor = torch.FloatTensor(radar_sim).unsqueeze(0) ecg_pred, _, _ = model(radar_tensor) ecg_pred_np = ecg_pred.squeeze().detach().numpy() ecg_gt_resampled = np.interp(np.linspace(0, 1, 256), t, ecg_gt) pcc, _ = pearsonr(ecg_pred_np, ecg_gt_resampled) print(f"ECG重建 PCC: {pcc:.2f}") assert pcc > 0.5, "PCC too low" print("✓ ECG quality test passed")
def test_noise_robustness(): """测试噪声鲁棒性""" model = test_forward_pass() t = np.linspace(0, 1, 1024) clean_signal = np.sin(2 * np.pi * 5 * t) for snr_db in [20, 10, 0, -5]: noise_power = np.var(clean_signal) / (10 ** (snr_db / 10)) noisy = clean_signal + np.random.randn(1024) * np.sqrt(noise_power) radar = torch.FloatTensor(noisy).unsqueeze(0).unsqueeze(0).expand(1, 8, 1024) ecg, _, _ = model(radar) signal_power = np.var(ecg.squeeze().detach().numpy()) print(f"SNR={snr_db:>3}dB | Output variance={signal_power:.4f}") print("✓ Noise robustness test passed")
def test_r_peak_detection(): """测试R峰检测""" model = test_forward_pass() t = np.linspace(0, 5, 2560) ecg = np.zeros_like(t) for i in range(5): center = 0.5 + i * 1.0 ecg += np.exp(-(t - center)**2 / 0.001) radar = np.tile(ecg, (8, 1))[:, :1024] radar += np.random.randn(8, 1024) * 0.1 radar_tensor = torch.FloatTensor(radar).unsqueeze(0) _, anchor_logits, _ = model(radar_tensor) anchor_probs = torch.sigmoid(anchor_logits).squeeze().detach().numpy() peaks, _ = find_peaks(anchor_probs, height=0.5, distance=100) print(f"检测到 {len(peaks)} 个R峰") assert len(peaks) >= 3, "R峰检测不足" print("✓ R-peak detection test passed")
if __name__ == "__main__": print("=" * 60) print("radarODE-MTL 测试套件") print("=" * 60) test_forward_pass() test_ecg_quality() test_noise_robustness() test_r_peak_detection() print("\n" + "=" * 60) print("所有测试通过 ✓") print("=" * 60)
|