From 4248e45d8ce6fc4e695587e1981b1f21ac3729bb Mon Sep 17 00:00:00 2001 From: Abuzar Mahmood Date: Wed, 5 Nov 2025 16:01:17 +0000 Subject: [PATCH] feat(visualization): add taste_names parameter to plotting functions Add optional taste_names parameter to all three main plotting functions: - generate_raster_plot() - generate_detailed_plot() - plot_env_pred_overlay() This allows users to specify custom taste names (e.g., 'Sucrose', 'NaCl') instead of generic labels like 'Taste 0', 'Taste 1'. Maintains backward compatibility by defaulting to 'Taste N' format when taste_names is not provided. Resolves #37 Co-authored-by: Ona --- src/visualize.py | 41 ++++++++++++++++++++++++++++++++--------- 1 file changed, 32 insertions(+), 9 deletions(-) diff --git a/src/visualize.py b/src/visualize.py index 7fca8c9..307dcf4 100644 --- a/src/visualize.py +++ b/src/visualize.py @@ -118,6 +118,7 @@ def plot_raster( def generate_raster_plot( segments_frame: pd.DataFrame, session_name: str = None, + taste_names: list = None, ): """ Generate a raster plot of the segments frame @@ -131,6 +132,8 @@ def generate_raster_plot( DataFrame containing the segments with classification results session_name : str Name of the session to display in the plot title + taste_names : list, optional + List of taste names corresponding to taste indices. If None, uses "Taste N" labels. Outputs: fig, ax : matplotlib figure and axis objects @@ -154,8 +157,12 @@ def generate_raster_plot( ax = [ax] for taste in range(len(taste_pred_array_list)): - ax[taste], im = plot_raster(taste_pred_array_list[taste], ax = ax[taste]) - ax[taste].set_ylabel(f'Taste {taste}\nTrial #') + ax[taste], im = plot_raster(taste_pred_array_list[taste], ax = ax[taste]) + if taste_names is not None and taste < len(taste_names): + taste_label = taste_names[taste] + else: + taste_label = f'Taste {taste}' + ax[taste].set_ylabel(f'{taste_label}\nTrial #') ax[0].set_title('Movement Classification') ax[-1].set_xlabel('Time (ms)') @@ -169,7 +176,7 @@ def generate_raster_plot( fig.suptitle(session_name) return fig, ax -def generate_detailed_plot(segments_frame, raw_emg=None, trial_idx=0, taste_idx=0): +def generate_detailed_plot(segments_frame, raw_emg=None, trial_idx=0, taste_idx=0, taste_names=None): """ Generate a detailed plot showing raw EMG signal and movement classifications. @@ -182,6 +189,7 @@ def generate_detailed_plot(segments_frame, raw_emg=None, trial_idx=0, taste_idx= raw_emg (np.ndarray, optional): Raw EMG signal data. If None, only shows classifications. trial_idx (int, optional): Trial index to visualize. Defaults to 0. taste_idx (int, optional): Taste index to visualize. Defaults to 0. + taste_names (list, optional): List of taste names corresponding to taste indices. If None, uses "Taste N" labels. Returns: tuple: (fig, ax) matplotlib figure and axes objects @@ -240,7 +248,11 @@ def generate_detailed_plot(segments_frame, raw_emg=None, trial_idx=0, taste_idx= # Set labels and title class_ax.set_ylabel('Movement Type') - class_ax.set_title(f'Taste {taste_idx}, Trial {trial_idx} - Movement Classifications') + if taste_names is not None and taste_idx < len(taste_names): + taste_label = taste_names[taste_idx] + else: + taste_label = f'Taste {taste_idx}' + class_ax.set_title(f'{taste_label}, Trial {trial_idx} - Movement Classifications') if raw_emg is None: class_ax.set_xlabel('Time (ms)') @@ -253,6 +265,7 @@ def plot_env_pred_overlay( raw_emg, cmap = None, mad_scale = 5, + taste_names = None, ): """ Create a comprehensive grid plot showing raw EMG signals with overlaid predictions. @@ -278,6 +291,7 @@ def plot_env_pred_overlay( - Blue (#3B75AF): MTMs Defaults to None. mad_scale (float, optional): Scaling factor for median absolute deviation + taste_names (list, optional): List of taste names corresponding to taste indices. If None, uses "Taste N" labels. Returns: tuple: Contains two elements: @@ -326,7 +340,11 @@ def plot_env_pred_overlay( if trial == raw_emg.shape[1] - 1: this_ax.set_xlabel('Time (ms)') if trial == 0: - this_ax.set_title(f'Taste {taste}') + if taste_names is not None and taste < len(taste_names): + taste_label = taste_names[taste] + else: + taste_label = f'Taste {taste}' + this_ax.set_title(taste_label) # Overlay predictions for _, segment in segments_frame.iterrows(): if segment.taste == taste and segment.trial == trial: @@ -383,17 +401,22 @@ def plot_env_pred_overlay( } cmap = ListedColormap(list(event_color_map.values()), name = 'NBT_cmap') - # Generate the plot + # Example taste names + taste_names = ['Sucrose', 'NaCl'] + + # Generate the plot with taste names fig, ax = generate_raster_plot( segments_frame = sample_frame, - session_name = "Example Visualization" + session_name = "Example Visualization", + taste_names = taste_names ) plt.show() - # Example of detailed plot + # Example of detailed plot with taste names fig, ax = generate_detailed_plot( segments_frame = sample_frame, trial_idx = 0, - taste_idx = 0 + taste_idx = 0, + taste_names = taste_names ) plt.show()