diff --git a/src/__main__.py b/src/__main__.py index acbbbf8..ca41293 100644 --- a/src/__main__.py +++ b/src/__main__.py @@ -1,5 +1,5 @@ import logging -import multiprocessing as mp +import torch.multiprocessing as mp from typing import NamedTuple from . import ingest @@ -14,35 +14,42 @@ logging.basicConfig( ) mp.set_start_method("spawn") +mp.log_to_stderr(logging.INFO) class Receiver(NamedTuple): ip: ingest.Host receiver: ingest.FeitReceiver + queue: "mp.Queue[ingest.CSI]" +manager = mp.Manager() receivers = [ - Receiver(ip, ingest.FeitReceiver(ip, mp.Queue(config.SAMPLE_RATE))) + Receiver(ip, ingest.FeitReceiver(ip), manager.Queue(config.SAMPLE_RATE)) for ip in config.RECEIVE_HOSTS ] # Start injecting CSI frames transmitter = ingest.FeitTransmitter() -webapp_queue: "mp.Queue[aoa.AoA]" = mp.Queue(config.SAMPLE_RATE) +webapp_queue: "mp.Queue[aoa.AoA]" = manager.Queue(config.SAMPLE_RATE) -processor = ingest.CSIProcessor( - {r.ip: r.receiver.processing_queue for r in receivers}, webapp_queue -) +visualise.aoa_queue = webapp_queue # Start webapp in background process -webapp = mp.Process(target=visualise.start, args=(webapp_queue,)) +webapp = mp.Process(target=visualise.start) webapp.start() -receiver_processes = [mp.Process(target=r.receiver.listen) for r in receivers] +receiver_processes = [ + mp.Process(target=r.receiver.listen, args=(r.queue,)) for r in receivers +] for proc in receiver_processes: proc.start() -processing_thread = mp.Process(target=processor.process_forever) +processing_thread = mp.Process( + target=ingest.CSIProcessor.process_forever, + args=({r.ip: r.queue for r in receivers}, webapp_queue), +) processing_thread.start() +processing_thread.join() diff --git a/src/ingest.py b/src/ingest.py index c55eac0..4247dcd 100644 --- a/src/ingest.py +++ b/src/ingest.py @@ -33,7 +33,7 @@ class FeitTransmitter: class FeitReceiver: - def __init__(self, host: Host, processing_queue: "mp.Queue[CSI]"): + def __init__(self, host: Host): self.host = host self.server = socket.socket(socket.AF_INET, socket.SOCK_DGRAM) self.server.connect(host) @@ -48,9 +48,8 @@ class FeitReceiver: self.logger = logging.getLogger( f"{__name__}.{self.__class__.__name__}-{self.host}" ) - self.processing_queue = processing_queue - def listen(self): + def listen(self, queue: "mp.Queue[CSI]"): prev_time = datetime.now() while True: # This is the max size of a UDP packet. The size of the actual CSI @@ -63,83 +62,71 @@ class FeitReceiver: f"Received CSI data after {datetime.now() - prev_time}" ) prev_time = datetime.now() - self.processing_queue.put(csidata) + queue.put(csidata) except struct.error: self.logger.error("Failed to parse CSI data") class CSIProcessor: - def __init__( - self, - receiver_connections: dict[Host, "mp.Queue[CSI]"], - webserver: "mp.Queue[AoA]", - ): - self.pending_data: dict[Host, tuple[datetime, CSI]] = {} - self.pending_data_lock = mp.Lock() - self.logger = logging.getLogger(f"{__name__}.{self.__class__.__name__}") - self.last_processed = datetime.now() - + def __init__(self): self.preprocess = Preprocessor() self.aoa = AoA() - self.connections = receiver_connections - self.webserver = webserver + @staticmethod + def process_data( + data: dict[Host, tuple[datetime, CSI]], + webserver: "mp.Queue[AoA]", + preprocess: Preprocessor, + aoa: AoA, + logger: logging.Logger, + ): + # Useful for figuring out the correct antenna order - RSSI values will decrease + # when the specific antenna is disconnected + rssis = [ + (ip, csi.header.rssi1, csi.header.rssi2) + for ip, (_, csi) in sorted(data.items()) + ] + logger.debug("Antenna RSSI values: {}".format(rssis)) - def add_data(self, host: Host, data: CSI): - if ( - host in self.pending_data - and self.last_processed < self.pending_data[host][0] - ): - self.logger.warning( - f"Skipping data from {host} at {self.pending_data[host][0]}" - ) - - with self.pending_data_lock: - self.pending_data[host] = (datetime.now(), data) - - # Useful for figuring out the correct antenna order - RSSI values will decrease - # when the specific antenna is disconnected - rssis = [ - (ip, csi.header.rssi1, csi.header.rssi2) - for ip, (_, csi) in sorted(self.pending_data.items()) - ] - self.logger.debug("Antenna RSSI values: {}".format(rssis)) - - def is_ready(self): - for host in config.RECEIVE_HOSTS: - if ( - host not in self.pending_data - or self.pending_data[host][0] <= self.last_processed - ): - return False - return True - - def process_data(self): - self.last_processed = datetime.now() antenna_data = [ - np.expand_dims(self.pending_data[ip][1].matrix[:, antenna], axis=2) + np.expand_dims(data[ip][1].matrix[:, antenna], axis=2) for ip, antenna in config.ANTENNA_ORDER ] # We have data from all servers all_data = np.concat(antenna_data, axis=1) - self.logger.info(f"Got final CSI data with shape {all_data.shape}") + logger.info(f"Got final CSI data with shape {all_data.shape}") - processed = self.preprocess.preprocess(all_data) + processed = preprocess.preprocess(all_data) processed_tensor = torch.tensor(processed, device=device) # visualise.add_data(all_data_tensor, processed) - self.aoa.update(processed_tensor) - self.logger.info("Processed data") - if not self.webserver.full(): - self.webserver.put(self.aoa) + aoa.update(processed_tensor) + logger.info("Processed data") + if not webserver.full(): + webserver.put(aoa) + + @staticmethod + def process_forever( + connections: dict[Host, "mp.Queue[CSI]"], webserver: "mp.Queue[AoA]" + ): + latest_data: dict[Host, tuple[datetime, CSI]] = {} + preprocessor = Preprocessor() + aoa = AoA() + + logger = mp.get_logger() + logger.info("Starting processing loop") - def process_forever(self): while True: - for ip, queue in self.connections.items(): + for ip, queue in connections.items(): while not queue.empty(): - self.add_data(ip, queue.get()) - if self.is_ready(): - self.process_data() + latest_data[ip] = (datetime.now(), queue.get()) + + sample_ready = all(ip in latest_data for ip in connections) + if sample_ready: + CSIProcessor.process_data( + latest_data, webserver, preprocessor, aoa, logger + ) + latest_data = {} else: - self.logger.debug("Not all data is ready") + logger.debug("Not all data is ready") time.sleep(0.001) diff --git a/src/visualise/__init__.py b/src/visualise/__init__.py index 1c27067..aaa5722 100644 --- a/src/visualise/__init__.py +++ b/src/visualise/__init__.py @@ -141,8 +141,5 @@ def aoa_tof(): ) -def start(conn: "mp.Queue[AoA]"): - global aoa_queue - - aoa_queue = conn +def start(): app.run(debug=True, use_reloader=False, host="0.0.0.0")