Files
templates/python/tests/test_spectrum.py

121 lines
6.1 KiB
Python

import cmath
import ctypes
import math
import os
import unittest
from protocan.spectrum import Filter, NativeSpectrum, Window
@unittest.skipUnless(os.environ.get("SETPROTOCOL_LIBRARY"), "Set SETPROTOCOL_LIBRARY to the built C library")
class SpectrumTests(unittest.TestCase):
@classmethod
def setUpClass(cls):
cls.core = NativeSpectrum(ctypes.CDLL(os.environ["SETPROTOCOL_LIBRARY"]))
def sample(self, n=4096, fs=1024, frequencies=(64,), amplitude=1):
times = [i / fs for i in range(n)]
return times, [amplitude * sum(math.sin(2 * math.pi * f * t) for f in frequencies) for t in times]
def test_fft_matches_independent_direct_dft(self):
times, _ = self.sample(32)
values = [math.sin(i * 1.37) + 0.1 * i for i in range(32)]
result = self.core.analyze(times, values, window=Window.RECT, remove_mean=False)
for k, amplitude in enumerate(result.amplitudes):
direct = abs(sum(v * cmath.exp(-2j * math.pi * k * i / 32) for i, v in enumerate(values))) / 32
if k not in (0, 16):
direct *= 2
self.assertAlmostEqual(direct, amplitude, places=11)
def test_all_windows_preserve_bin_centered_peak_amplitude(self):
times, values = self.sample(amplitude=3.25)
for window in Window:
with self.subTest(window=window):
result = self.core.analyze(times, values, window=window)
peak = max(range(len(result.amplitudes)), key=result.amplitudes.__getitem__)
self.assertEqual(64, result.frequencies[peak])
self.assertAlmostEqual(3.25, result.amplitudes[peak], places=9)
detected = self.core.dominant_peak(result)
self.assertAlmostEqual(64, detected.frequency_hz, places=9)
self.assertAlmostEqual(3.25, detected.amplitude, places=9)
def test_dc_and_nyquist_are_not_doubled(self):
times, _ = self.sample()
result = self.core.analyze(times, [2.5] * len(times), remove_mean=False, window=Window.RECT)
self.assertAlmostEqual(2.5, result.amplitudes[0], places=10)
result = self.core.analyze(times, [3 * (-1) ** i for i in range(len(times))], window=Window.RECT)
self.assertAlmostEqual(3, result.amplitudes[-1], places=10)
result = self.core.analyze(times, [2.5] * len(times))
self.assertLess(max(result.amplitudes), 1e-12)
def test_windows_suppress_far_leakage_and_flattop_recovers_off_bin_amplitude(self):
times, values = self.sample(frequencies=(64.13,))
rect = self.core.analyze(times, values, window=Window.RECT)
hann = self.core.analyze(times, values, window=Window.HANN)
flat = self.core.analyze(times, values, window=Window.FLATTOP)
self.assertLess(hann.amplitudes[400], rect.amplitudes[400] / 100)
self.assertAlmostEqual(1, max(flat.amplitudes), delta=0.002)
def test_filters_attenuate_expected_bands(self):
times, values = self.sample(frequencies=(16, 64, 256))
low = self.core.analyze(times, values, filter=Filter.LOW_PASS, high_hz=64)
high = self.core.analyze(times, values, filter=Filter.HIGH_PASS, low_hz=64)
band = self.core.analyze(times, values, filter=Filter.BAND_PASS, low_hz=32, high_hz=128)
notch = self.core.analyze(times, values, filter=Filter.NOTCH, low_hz=64)
at = lambda result, hz: result.amplitudes[int(hz / (result.sample_rate / result.size))]
self.assertGreater(at(low, 16), 0.99)
self.assertLess(at(low, 256), 0.05)
self.assertAlmostEqual(1 / math.sqrt(2), at(low, 64), delta=0.001)
self.assertLess(at(high, 16), 0.07)
self.assertGreater(at(high, 256), 0.99)
self.assertGreater(at(band, 64), 0.93)
self.assertLess(at(band, 16), 0.25)
self.assertLess(at(band, 256), 0.2)
self.assertLess(at(notch, 64), 0.02)
self.assertGreater(at(notch, 16), 0.99)
def test_timestamp_rate_not_requested_rate_and_jitter_interpolation(self):
n, fs = 1024, 200
times = [(i + (0.05 if i % 2 else 0)) / fs for i in range(n)]
values = [2 * math.sin(2 * math.pi * (16 * fs / n) * t) for t in times]
result = self.core.analyze(times, values)
self.assertAlmostEqual((n - 1) / (times[-1] - times[0]), result.sample_rate)
self.assertGreater(result.jitter, 0.04)
self.assertAlmostEqual(2, result.amplitudes[16], delta=0.005)
def test_rejects_gaps_duplicates_bad_values_and_cutoffs(self):
times, values = self.sample(128)
for broken in ([0.0] * 128, times[:64] + [t + 1 for t in times[64:]], list(reversed(times))):
with self.assertRaisesRegex(ValueError, "timestamps"):
self.core.analyze(broken, values)
for value in (float("nan"), float("inf")):
with self.assertRaises(ValueError):
self.core.analyze(times, values[:-1] + [value])
for cutoff in (0, -1, 512, 1000, float("nan")):
with self.assertRaisesRegex(ValueError, "Filter frequencies"):
self.core.analyze(times, values, filter=Filter.LOW_PASS, high_hz=cutoff)
with self.assertRaises(ValueError):
self.core.analyze(times, values, filter=Filter.BAND_PASS, low_hz=100, high_hz=50)
def test_size_limits_tail_selection_and_inputs_unchanged(self):
times, values = self.sample(1000)
original = values[:]
result = self.core.analyze(times, values)
self.assertEqual(512, result.size)
self.assertEqual(original, values)
self.assertEqual(256, self.core.analyze(times, values, max_size=256).size)
for n in (0, 1, 15):
with self.assertRaisesRegex(ValueError, "16 samples"):
self.core.analyze(times[:n], values[:n])
for size in (0, 15, 1000, 32768):
with self.assertRaises(ValueError):
self.core.analyze(times, values, max_size=size)
with self.assertRaises(ValueError):
self.core.analyze(times[:-1], values)
times, values = self.sample(20000)
self.assertEqual(16384, self.core.analyze(times, values, max_size=16384).size)
if __name__ == "__main__":
unittest.main()