import logging from queue import Queue from typing import Any, Callable import numpy as np import numpy.typing as npt from scipy.signal import butter, correlate, sosfilt, sosfilt_zi from where_fi.collection import CSIMatrix from where_fi.collection.csi_frame import CSI from where_fi.collection.protocols import CSIHost from where_fi.config import config from ..visualise import server as visualise logger = logging.getLogger(__name__) np.seterr(invalid="ignore") class Preprocessor: def __init__(self) -> None: self.short_term_avg = np.zeros((1,), dtype=np.complex64) self.long_term_avg = np.zeros((1,), dtype=np.complex64) # self.filter = butter( # 5, # config.preprocessing.bandpass.bounds, # fs=config.sample_rate, # btype="band", # output="sos", # ) self._last_sample = None self.denoising_samples = int( config.preprocessing.denoising_period * config.collection_sample_rate ) self.circular_buffer = np.zeros( ( self.denoising_samples, # Number of samples config.subcarriers, # Number of subcarriers config.antennas.count, # Number of RX antennas 1, # Number of TX antennas ), dtype=np.complex64, ) self.sample_index = 0 @property def last_sample(self) -> None | CSIMatrix: """ The last sample of the preprocessor. This is used for low frequency processing """ match config.preprocessing.denoising: case "none": return self._last_sample case "median": # Return the median of the last samples median_abs = np.median(np.abs(self.circular_buffer), axis=0) median_angle = np.median(np.angle(self.circular_buffer), axis=0) ans = median_abs * np.exp(1j * median_angle) return ans case "mean": # Return the mean of the last samples return np.mean(self.circular_buffer, axis=0) case _: raise ValueError( f"Invalid denoising method: {config.preprocessing.denoising}" ) def remove_sto(self, csi: CSIMatrix) -> CSIMatrix: """ Remove sampling time offsets caused by: - Sampling frequency offset - Packet detection delay This is done by multiplying the CSI matrices of consecutive antennas in the array. According to [1]: > Conjugate multiplication and division are the only two methods to > eliminate the SFO and PDD. No citation or explanation is provided, so not sure why/whether it works. Something similar is also done in [2] without explanation. [1] - https://tns.thss.tsinghua.edu.cn/wst/docs/sanitization [2] - https://doi.org/10.1109/ICC51166.2024.10623053 """ csi_remove_sto = np.zeros_like(csi) for antenna in range(csi.shape[1]): antenna_nxt = (antenna + 1) % csi.shape[1] csi_remove_sto[:, antenna, :] = np.multiply( csi[:, antenna, :], csi[:, antenna_nxt, :].conj() ) return csi_remove_sto def remove_sfo( self, csi: CSIMatrix, visualiser: None | Callable[[npt.NDArray[Any], visualise.figures.Figure], None] = None, ) -> CSIMatrix: """ Remove sampling frequency offsets by linear regression. This is caused by the difference in sampling frequency between the transmitter and receiver, and this is linear in frequency. We can estimate it using linear regression and compensate for it. """ N_st, N_rx, _ = csi.shape unwrapped = np.unwrap(np.angle(csi[:, :, 0]), axis=0).reshape(N_st, N_rx, 1) if visualiser: visualiser(unwrapped, visualise.figures.Figure.UNWRAPPED_PHASE) X = np.vstack([np.arange(N_st), np.ones(N_st)]).T for antenna in range(csi.shape[1]): tau, rho = np.linalg.lstsq(X, unwrapped[:, antenna, 0])[0] csi[:, antenna, 0] = np.abs(csi[:, antenna, 0]) * np.exp( 1j * (np.angle(csi[:, antenna, 0]) - (tau * np.arange(csi.shape[0]) + rho)) ) return csi def normalise_magnitude( self, csi: CSIMatrix, frames: dict[CSIHost, CSI] ) -> CSIMatrix: """ Normalise the magnitude of the CSI data to compensate for the effect of the Automatic Gain Control (AGC) of the receiver References: [1] - """ rssi = [ frames[host].header.rssi1 if antenna == 0 else frames[host].header.rssi2 for host, antenna in config.antennas.order ] rssi_linear = np.reshape(10 ** (np.array(rssi) / 10), (1, -1, 1)) csi_power = np.sum(np.abs(csi) ** 2) return csi * np.sqrt(rssi_linear / csi_power) def preprocess( self, h: CSIMatrix, frames: dict[CSIHost, CSI], visualiser: None | Callable[[npt.NDArray[Any], visualise.figures.Figure], None] = None, ) -> CSIMatrix: # CSI data is not available for pilot subcarriers. h_hat: CSIMatrix = np.where( np.expand_dims(h[:, 0, 0] == 0, axis=(1, 2)), correlate(h, [[[1 / 2]], [[0]], [[1 / 2]]], mode="same"), h, ) logger.debug(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") h_hat = self.normalise_magnitude(h_hat, frames) h_hat = self.remove_sfo(h_hat, visualiser=visualiser) # Skip subcarrierss per config h_hat = h_hat[:: config.preprocessing.subcarrier_step, :, :] # h_hat = self.remove_sto(h_hat) # 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.complex64) 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 self._last_sample = h_hat if config.preprocessing.denoising != "none": self.circular_buffer[self.sample_index] = h_hat self.sample_index = (self.sample_index + 1) % self.denoising_samples return h_hat # 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]