From af9faa3d078c4849badbb04ac54114b5362689bc Mon Sep 17 00:00:00 2001 From: Christos Falas Date: Wed, 14 May 2025 16:48:08 +0100 Subject: [PATCH] add support for custom visualisations --- where_fi/application.py | 16 ++++++++++++++-- where_fi/visualise/server/__init__.py | 8 ++++---- where_fi/visualise/server/figures.py | 6 +++++- 3 files changed, 23 insertions(+), 7 deletions(-) diff --git a/where_fi/application.py b/where_fi/application.py index 2144e3e..d51400a 100644 --- a/where_fi/application.py +++ b/where_fi/application.py @@ -12,6 +12,7 @@ from where_fi.collection.protocols import CSIProducer, MergedCSI from where_fi.config import config from where_fi.processing.preprocess import Preprocessor from where_fi.visualise import server as visualise +from where_fi.visualise.server import figures class Receiver(NamedTuple): @@ -89,9 +90,20 @@ class CSIApplication: self.processing_callback = None self.preprocessor = Preprocessor() + self.custom_figures: dict[ + visualise.figures.FigureId, visualise.figures.SpecificFigure + ] = {} + + def register_figure( + self, + figure_id: visualise.figures.FigureId, + figure: visualise.figures.SpecificFigure, + ) -> None: + self.custom_figures[figure_id] = figure + figures.all_figures[figure_id] = figure def visualise_data( - self, data: npt.NDArray[Any], dtype: visualise.figures.Figure + self, data: npt.NDArray[Any], dtype: visualise.figures.FigureId ) -> None: """ Update a visualisation with the given data. The available visualisations are as @@ -155,7 +167,7 @@ class CSIApplication: self.raw_csi_callback(sample.matrix) if self.pre_merge_callback is not None: self.pre_merge_callback(sample.frames) - processed = self.preprocessor.preprocess(sample.matrix) + processed = self.preprocessor.preprocess(sample.matrix, sample.frames) if self.visualise_raw: self.visualise_data(processed, visualise.figures.Figure.PROCESSED_CSI) if self.preprocessed_csi_callback is not None: diff --git a/where_fi/visualise/server/__init__.py b/where_fi/visualise/server/__init__.py index 9437fd8..f1fa719 100644 --- a/where_fi/visualise/server/__init__.py +++ b/where_fi/visualise/server/__init__.py @@ -18,7 +18,7 @@ from . import figures @dataclass class VisualiserData: data: npt.NDArray[Any] - dtype: figures.Figure + dtype: figures.FigureId class FigureServer(figure_pb2_grpc.FigureServiceServicer): @@ -30,8 +30,8 @@ class FigureServer(figure_pb2_grpc.FigureServiceServicer): def GetFigure( self, request: figure_pb2.FigureRequest, context: grpc.ServicerContext ) -> Generator[figure_pb2.Figure, None, None]: - for figure_group in figures.Figure: - for figure in figure_group.figure_class().figures: + for figure_group in figures.all_figures.values(): + for figure in figure_group.figures: yield figure def GetFigureUpdate( @@ -65,7 +65,7 @@ class Webapp: self.figure_server = FigureServer() def add_data( - self, dtype: figures.Figure, new_data: npt.NDArray[np.complex128] + self, dtype: figures.FigureId, new_data: npt.NDArray[np.complex128] ) -> None: updates = figures.all_figures[dtype].update(new_data) for fig_id, update in updates.items(): diff --git a/where_fi/visualise/server/figures.py b/where_fi/visualise/server/figures.py index 14103dc..5097acf 100644 --- a/where_fi/visualise/server/figures.py +++ b/where_fi/visualise/server/figures.py @@ -271,12 +271,15 @@ class Figure(Enum): MUSIC_EIGENVALUES = 3 AOA_HEATMAP = 4 PHASE_ANALYSIS = 5 + MAGN_ANALYSIS = 6 def figure_class(self) -> SpecificFigure: return all_figures[self] -all_figures = { +FigureId = str | Figure + +all_figures: dict[FigureId, SpecificFigure] = { Figure.RAW_CSI: PerAntennaFigure( [ SimpleLineChart("Raw CSI Phase", "Subcarrier", "Phase"), @@ -297,4 +300,5 @@ all_figures = { Figure.MUSIC_EIGENVALUES: MusicEigenvalueHistogram(), Figure.AOA_HEATMAP: HeatmapFigure(), Figure.PHASE_ANALYSIS: RandomVariable("Phase"), + Figure.MAGN_ANALYSIS: RandomVariable("Magnitude"), }