Fix processes without forking for CUDA
This commit is contained in:
parent
5f6824378c
commit
efe8923c51
@ -1,5 +1,5 @@
|
|||||||
import logging
|
import logging
|
||||||
import multiprocessing as mp
|
import torch.multiprocessing as mp
|
||||||
from typing import NamedTuple
|
from typing import NamedTuple
|
||||||
|
|
||||||
from . import ingest
|
from . import ingest
|
||||||
@ -14,35 +14,42 @@ logging.basicConfig(
|
|||||||
)
|
)
|
||||||
|
|
||||||
mp.set_start_method("spawn")
|
mp.set_start_method("spawn")
|
||||||
|
mp.log_to_stderr(logging.INFO)
|
||||||
|
|
||||||
|
|
||||||
class Receiver(NamedTuple):
|
class Receiver(NamedTuple):
|
||||||
ip: ingest.Host
|
ip: ingest.Host
|
||||||
receiver: ingest.FeitReceiver
|
receiver: ingest.FeitReceiver
|
||||||
|
queue: "mp.Queue[ingest.CSI]"
|
||||||
|
|
||||||
|
|
||||||
|
manager = mp.Manager()
|
||||||
receivers = [
|
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
|
for ip in config.RECEIVE_HOSTS
|
||||||
]
|
]
|
||||||
|
|
||||||
# Start injecting CSI frames
|
# Start injecting CSI frames
|
||||||
transmitter = ingest.FeitTransmitter()
|
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(
|
visualise.aoa_queue = webapp_queue
|
||||||
{r.ip: r.receiver.processing_queue for r in receivers}, webapp_queue
|
|
||||||
)
|
|
||||||
|
|
||||||
# Start webapp in background process
|
# Start webapp in background process
|
||||||
webapp = mp.Process(target=visualise.start, args=(webapp_queue,))
|
webapp = mp.Process(target=visualise.start)
|
||||||
webapp.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:
|
for proc in receiver_processes:
|
||||||
proc.start()
|
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.start()
|
||||||
|
processing_thread.join()
|
||||||
|
|||||||
@ -33,7 +33,7 @@ class FeitTransmitter:
|
|||||||
|
|
||||||
|
|
||||||
class FeitReceiver:
|
class FeitReceiver:
|
||||||
def __init__(self, host: Host, processing_queue: "mp.Queue[CSI]"):
|
def __init__(self, host: Host):
|
||||||
self.host = host
|
self.host = host
|
||||||
self.server = socket.socket(socket.AF_INET, socket.SOCK_DGRAM)
|
self.server = socket.socket(socket.AF_INET, socket.SOCK_DGRAM)
|
||||||
self.server.connect(host)
|
self.server.connect(host)
|
||||||
@ -48,9 +48,8 @@ class FeitReceiver:
|
|||||||
self.logger = logging.getLogger(
|
self.logger = logging.getLogger(
|
||||||
f"{__name__}.{self.__class__.__name__}-{self.host}"
|
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()
|
prev_time = datetime.now()
|
||||||
while True:
|
while True:
|
||||||
# This is the max size of a UDP packet. The size of the actual CSI
|
# 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}"
|
f"Received CSI data after {datetime.now() - prev_time}"
|
||||||
)
|
)
|
||||||
prev_time = datetime.now()
|
prev_time = datetime.now()
|
||||||
self.processing_queue.put(csidata)
|
queue.put(csidata)
|
||||||
except struct.error:
|
except struct.error:
|
||||||
self.logger.error("Failed to parse CSI data")
|
self.logger.error("Failed to parse CSI data")
|
||||||
|
|
||||||
|
|
||||||
class CSIProcessor:
|
class CSIProcessor:
|
||||||
def __init__(
|
def __init__(self):
|
||||||
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()
|
|
||||||
|
|
||||||
self.preprocess = Preprocessor()
|
self.preprocess = Preprocessor()
|
||||||
self.aoa = AoA()
|
self.aoa = AoA()
|
||||||
|
|
||||||
self.connections = receiver_connections
|
@staticmethod
|
||||||
self.webserver = webserver
|
def process_data(
|
||||||
|
data: dict[Host, tuple[datetime, CSI]],
|
||||||
def add_data(self, host: Host, data: CSI):
|
webserver: "mp.Queue[AoA]",
|
||||||
if (
|
preprocess: Preprocessor,
|
||||||
host in self.pending_data
|
aoa: AoA,
|
||||||
and self.last_processed < self.pending_data[host][0]
|
logger: logging.Logger,
|
||||||
):
|
):
|
||||||
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
|
# Useful for figuring out the correct antenna order - RSSI values will decrease
|
||||||
# when the specific antenna is disconnected
|
# when the specific antenna is disconnected
|
||||||
rssis = [
|
rssis = [
|
||||||
(ip, csi.header.rssi1, csi.header.rssi2)
|
(ip, csi.header.rssi1, csi.header.rssi2)
|
||||||
for ip, (_, csi) in sorted(self.pending_data.items())
|
for ip, (_, csi) in sorted(data.items())
|
||||||
]
|
]
|
||||||
self.logger.debug("Antenna RSSI values: {}".format(rssis))
|
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 = [
|
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
|
for ip, antenna in config.ANTENNA_ORDER
|
||||||
]
|
]
|
||||||
|
|
||||||
# We have data from all servers
|
# We have data from all servers
|
||||||
all_data = np.concat(antenna_data, axis=1)
|
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)
|
processed_tensor = torch.tensor(processed, device=device)
|
||||||
# visualise.add_data(all_data_tensor, processed)
|
# visualise.add_data(all_data_tensor, processed)
|
||||||
self.aoa.update(processed_tensor)
|
aoa.update(processed_tensor)
|
||||||
self.logger.info("Processed data")
|
logger.info("Processed data")
|
||||||
if not self.webserver.full():
|
if not webserver.full():
|
||||||
self.webserver.put(self.aoa)
|
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:
|
while True:
|
||||||
for ip, queue in self.connections.items():
|
for ip, queue in connections.items():
|
||||||
while not queue.empty():
|
while not queue.empty():
|
||||||
self.add_data(ip, queue.get())
|
latest_data[ip] = (datetime.now(), queue.get())
|
||||||
if self.is_ready():
|
|
||||||
self.process_data()
|
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:
|
else:
|
||||||
self.logger.debug("Not all data is ready")
|
logger.debug("Not all data is ready")
|
||||||
time.sleep(0.001)
|
time.sleep(0.001)
|
||||||
|
|||||||
@ -141,8 +141,5 @@ def aoa_tof():
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
def start(conn: "mp.Queue[AoA]"):
|
def start():
|
||||||
global aoa_queue
|
|
||||||
|
|
||||||
aoa_queue = conn
|
|
||||||
app.run(debug=True, use_reloader=False, host="0.0.0.0")
|
app.run(debug=True, use_reloader=False, host="0.0.0.0")
|
||||||
|
|||||||
Loading…
Reference in New Issue
Block a user