import logging from queue import Queue import numpy as np import numpy.typing as npt from scipy.signal import butter, correlate, sosfilt, sosfilt_zi from ..config import config logger = logging.getLogger(__name__) logger.setLevel(logging.DEBUG) np.seterr(invalid="ignore") class Preprocessor: def __init__(self) -> None: self.prev_entries: Queue[npt.NDArray[np.complex128]] = Queue(maxsize=100) self.short_term_avg = np.zeros((1,), dtype=np.complex128) self.long_term_avg = np.zeros((1,), dtype=np.complex128) self.filter = butter( 5, config.preprocessing.bandpass.bounds, fs=config.sample_rate, btype="band", output="sos", ) def preprocess(self, h: npt.NDArray[np.complex128]) -> npt.NDArray[np.complex128]: # CSI data is not available for pilot subcarriers. h_hat = np.where( np.expand_dims(h[:, 0, 0] == 0, axis=(1, 2)), correlate(h, [[[1 / 2]], [[0]], [[1 / 2]]], mode="same"), h, ) # Skip every other subcarrier # h_hat = h_hat[::2, :, :] # logger.info(f"CSI shape: {h_hat.shape}") h_hat = np.multiply(h_hat, h_hat.conj() / abs(h_hat.conj())) h_hat = np.nan_to_num(h_hat) # h_hat = correlate(h_hat, np.ones((3, 1, 1)) / 3) h_hat = correlate(h_hat, [[[1 / 4]], [[1 / 2]], [[1 / 4]]], mode="valid") # Assume that all csi matrices will have the same shape if self.long_term_avg.shape != h_hat.shape: self.long_term_avg = np.zeros(h_hat.shape, dtype=np.complex128) self.long_term_avg = ( self.long_term_avg * (1 - config.preprocessing.moving_average_alpha) + h_hat * config.preprocessing.moving_average_alpha ) # Remove long term average, to remove static paths h_hat -= self.long_term_avg # Apply bandpass filter to remove low and high frequency noise if not hasattr(self, "filter_zi"): self.filter_zi = ( np.expand_dims(sosfilt_zi(self.filter), axis=(-1, -2, -3)) * h_hat ) h_hat_filt, self.filter_zi = sosfilt( self.filter, [h_hat], zi=self.filter_zi, axis=0 ) return h_hat_filt[0]