CoolFace
Apppublic

hanhou/patchseq

sourceHugging Faceupdated 7mo agoView on Hugging Face
0likes
patchseq_panel_viz.py1080 linesDownload Raw Back to code
1"""2Panel-based visualization tool for navigating and visualizing patch-seq NWB files.3 4To start the app, run:5panel serve panel_nwb_viz.py --dev --allow-websocket-origin=codeocean.allenneuraldynamics.org --title "Patch-seq Data Explorer"  # noqa: E5016"""7 8import logging9import re10 11import numpy as np12import pandas as pd13import panel as pn14import param15from bokeh.io import curdoc16from bokeh.layouts import column as bokeh_column17from bokeh.models import (18    BoxZoomTool,19)20from bokeh.plotting import figure21 22from LCNE_patchseq_analysis.data_util.metadata import load_ephys_metadata23from LCNE_patchseq_analysis.data_util.nwb import PatchSeqNWB24from LCNE_patchseq_analysis.pipeline_util.s3 import (25    S3_PUBLIC_URL_BASE,26    get_public_url_cell_summary,27    get_public_url_sweep,28    load_efel_features_from_roi,29)30from LCNE_patchseq_analysis.population_analysis.spikes import (31    extract_representative_spikes,32)33 34from components.scatter_plot import ScatterPlot35from components.spike_analysis import RawSpikeAnalysis36 37 38logging.basicConfig(level=logging.DEBUG)39logger = logging.getLogger(__name__)40 41# Initialize Panel with Bootstrap and Tabulator extensions42pn.extension("tabulator")43curdoc().title = "LC-NE Patch-seq Data Explorer"44 45 46def add_efel_asymmetry_columns(df_meta: pd.DataFrame) -> None:47    rise_cols = {}48    fall_cols = {}49 50    for col in df_meta.columns:51        if not isinstance(col, str):52            continue53        if not col.startswith("efel_") or " @ " not in col:54            continue55        metric, suffix = col.split(" @ ", 1)56        match = re.match(r"^(efel_.+?)_(rise|fall)(.*)$", metric)57        if not match:58            continue59        base = f"{match.group(1)}{match.group(3)}"60        key = f"{base} @ {suffix}"61        if match.group(2) == "rise":62            rise_cols[key] = col63        else:64            fall_cols[key] = col65 66    for key in sorted(set(rise_cols) & set(fall_cols)):67        base, suffix = key.split(" @ ", 1)68        new_col = f"{base}_asymmetry @ {suffix}"69        rise_vals = pd.to_numeric(df_meta[rise_cols[key]], errors="coerce")70        fall_vals = pd.to_numeric(df_meta[fall_cols[key]], errors="coerce")71        df_meta[new_col] = np.where(fall_vals != 0, rise_vals / fall_vals, np.nan)72 73 74class PatchSeqNWBApp(param.Parameterized):75    """76    Object-Oriented Panel App for navigating NWB files.77    Encapsulates metadata loading, sweep visualization, and cell selection.78    """79 80    class DataHolder(param.Parameterized):81        """82        Holder for currently selected cell ID and sweep number.83        """84 85        ephys_roi_id_selected = param.String(default="")86        sweep_number_selected = param.Integer(default=0)87        filtered_df_meta = param.DataFrame()88 89    def __init__(self):90        """91        Initialize the PatchSeqNWBApp.92        """93        # Holder for currently selected cell ID.94        self.data_holder = PatchSeqNWBApp.DataHolder()95 96        # Store figures and state for context-aware export97        self._cell_explorer_figures = {}98        self._active_tab = 099        self._export_progress = pn.widgets.Progress(100            name="Export progress",101            value=0,102            max=100,103            sizing_mode="stretch_width",104            bar_color="success",105            active=True,106            visible=False,107        )108 109        # Load and prepare metadata.110        self.df_meta = load_ephys_metadata(111            if_from_s3=True, if_with_seq=True, if_with_morphology=True112        )113        add_efel_asymmetry_columns(self.df_meta)114        self.df_meta.rename(115            columns={116                "x": "X (A --> P)",117                "y": "Y (D --> V)",118                "z": "Z (L --> R)",119            },120            inplace=True,121        )122 123        # Preprocess data: Convert NaN in "virus" column to "None"124        if "virus" in self.df_meta.columns:125            self.df_meta["virus"] = self.df_meta["virus"].fillna("None")126 127        self.cell_key = [128            "Date",129            "jem-id_cell_specimen",130            "ephys_roi_id",131            "ephys_qc",132            "LC_targeting",133            "injection region",134            "Y (D --> V)",135        ]136        # Turn Date to datetime137        self.df_meta.loc[:, "Date_str"] = self.df_meta[138            "Date"139        ]  # Keep the original Date as string140        self.df_meta.loc[:, "Date"] = pd.to_datetime(141            self.df_meta["Date"], errors="coerce"142        )143 144        # Initialize scatter plot component145        self.scatter_plot = ScatterPlot(self.df_meta, self.data_holder)146 147        # Initialize spike analysis component148        self.raw_spike_analysis = RawSpikeAnalysis(self.df_meta, main_app=self)149 150        # Create a copy for filtering - this will be updated by the global filter151        self.data_holder.filtered_df_meta = self.df_meta.copy()152 153    def update_bokeh(self, raw, sweep, downsample_factor=3):154        """155        Update the Bokeh plot for a given sweep.156        """157        trace = raw.get_raw_trace(sweep)[::downsample_factor]158        stimulus = raw.get_stimulus(sweep)[::downsample_factor]159        time = raw.get_time(sweep)[::downsample_factor]160 161        box_zoom_auto = BoxZoomTool(dimensions="auto")162 163        # Create the voltage trace plot164        voltage_plot = figure(165            title=f"Full traces - Sweep number {sweep} (downsampled {downsample_factor}x)",166            height=300,167            tools=["hover", box_zoom_auto, "box_zoom", "wheel_zoom", "reset", "pan"],168            active_drag=box_zoom_auto,169            x_range=(0, time[-1]),170            y_axis_label="Vm (mV)",171            sizing_mode="stretch_width",172        )173        voltage_plot.line(time, trace, line_width=1.5, color="navy")174 175        # Create the stimulus plot176        stim_plot = figure(177            height=150,178            tools=["hover", box_zoom_auto, "box_zoom", "wheel_zoom", "reset", "pan"],179            active_drag=box_zoom_auto,180            x_range=voltage_plot.x_range,  # Link x ranges181            x_axis_label="Time (ms)",182            y_axis_label="I (pA)",183            sizing_mode="stretch_width",184        )185        stim_plot.line(time, stimulus, line_width=1.5, color="firebrick")186 187        # Store figures for export188        self._cell_explorer_figures = {189            "voltage_trace": voltage_plot,190            "stimulus_trace": stim_plot,191        }192 193        # Stack the plots vertically using bokeh's column layout194        layout = bokeh_column(195            voltage_plot, stim_plot, sizing_mode="stretch_width", margin=(50, 0, 0, 0)196        )197        return layout198 199    @staticmethod200    def highlight_selected_rows(row, highlight_subset, color, fields=None):201        """202        Highlight rows based on a subset of values.203        If fields is None, highlight the entire row.204        """205        style = [""] * len(row)206        if row["sweep_number"] in highlight_subset:207            if fields is None:208                return [f"background-color: {color}"] * len(row)209            else:210                for field in fields:211                    style[list(row.keys()).index(field)] = f"background-color: {color}"212        return style213 214    @staticmethod215    def get_qc_message(sweep, df_sweeps):216        """Return a QC message based on sweep data."""217        if sweep not in df_sweeps["sweep_number"].values:218            return "<span style='color:red;'>Invalid sweep!</span>"219        if sweep in df_sweeps.query("passed != passed")["sweep_number"].values:220            return "<span style='background:salmon;'>Sweep terminated by the experimenter!</span>"221        if sweep in df_sweeps.query("passed == False")["sweep_number"].values:222            return (223                f"<span style='background:yellow;'>Sweep failed QC! "224                f"({df_sweeps[df_sweeps.sweep_number == sweep].reasons.iloc[0][0]})</span>"225            )226        return "<span style='background:lightgreen;'>Sweep passed QC!</span>"227 228    def _download_svg_context_aware(self):229        """Download SVGs based on the active tab."""230        from components.utils.svg_export import export_figures_to_svg_zip231 232        progress = self._export_progress233        progress.value = 0234        progress.visible = True235        progress.name = "Preparing export..."236 237        try:238            if self._active_tab == 0:  # Cell Explorer tab239                figures = {}240                progress.name = "Collecting scatter figures..."241 242                if (243                    hasattr(self.scatter_plot, "_latest_figures")244                    and self.scatter_plot._latest_figures245                ):246                    figures.update(self.scatter_plot._latest_figures)247 248                progress.name = "Collecting cell explorer figures..."249 250                if self._cell_explorer_figures:251                    figures.update(self._cell_explorer_figures)252 253                if not figures:254                    raise RuntimeError(255                        "No figures available. Select a cell and generate scatter plot first."256                    )257 258                prefix = "app_download_cell_explorer"259            elif self._active_tab == 1:  # Spike Analysis tab260                progress.name = "Collecting spike analysis figures..."261 262                if not self.raw_spike_analysis._latest_figures:263                    raise RuntimeError(264                        "Generate the spike analysis plots before downloading SVGs."265                    )266                figures = self.raw_spike_analysis._latest_figures267                prefix = "app_download_spike_analysis"268            else:  # Other tabs269                raise RuntimeError(270                    f"SVG export not supported for tab {self._active_tab}."271                )272 273            exportable_count = sum(1 for fig in figures.values() if fig is not None)274            if exportable_count == 0:275                raise RuntimeError(276                    "No exportable figures available for the current tab."277                )278 279            progress.value = 0280            progress.name = f"Rendering SVGs (0/{exportable_count})"281 282            def report_progress(current, total):283                progress.value = min(100, int((current / total) * 100))284                progress.name = f"Rendering SVGs ({current}/{total})"285 286            zip_buffer, timestamp = export_figures_to_svg_zip(287                figures, progress_callback=report_progress288            )289 290            progress.value = 100291            progress.name = "Packing download..."292            self._download_button.filename = f"{prefix}_{timestamp}.zip"293            progress.name = "Ready"294 295            return zip_buffer296        finally:297            progress.visible = False298            progress.value = 0299            progress.name = "Export progress"300 301    def apply_global_filter(self, query_string):302        """303        Apply a query filter to the metadata DataFrame.304 305        Args:306            query_string: A string in pandas query format to filter the metadata307 308        Returns:309            Filtered DataFrame310        """311        if not query_string.strip():312            # If query is empty, reset to the full dataset313            self.data_holder.filtered_df_meta = self.df_meta.copy()314            return f"Reset to full dataset (N={len(self.data_holder.filtered_df_meta)})"315 316        try:317            # Apply the filter query318            filtered = self.df_meta.query(query_string)319            if len(filtered) == 0:320                return "Query returned 0 results. Filter not applied."321 322            # Update the filtered dataframe in the data_holder323            # This will trigger updates in any components bound to this parameter324            self.data_holder.filtered_df_meta = filtered325 326            return f"Query applied. {len(filtered)} records match (out of {len(self.df_meta)})."327        except Exception as e:328            return f"Error in query: {str(e)}"329 330    def create_scatter_plot(self):331        """332        Create the scatter plot panel using the ScatterPlot component.333        """334 335        # Get plot controls from the scatter plot component336        controls = self.scatter_plot.controls337        control_width = 300338 339        # Create a reactive scatter plot that updates when controls change340        scatter_plot = pn.bind(341            self.scatter_plot.update_scatter_plot,342            controls["x_axis_select"].param.value,343            controls["y_axis_select"].param.value,344            controls["color_col_select"].param.value,345            controls["color_palette_select"].param.value,346            controls["size_col_select"].param.value,347            controls["size_range_slider"].param.value,348            controls["size_gamma_slider"].param.value,349            controls["alpha_slider"].param.value,350            controls["width_slider"].param.value,351            controls["height_slider"].param.value,352            controls["font_size_slider"].param.value,353            controls["bins_slider"].param.value,354            controls["hist_height_slider"].param.value,355            controls["show_gmm"].param.value,356            controls["n_components_x"].param.value,357            controls["n_components_y"].param.value,358            controls["show_linear_fit"].param.value,359            df_meta=self.data_holder.param.filtered_df_meta,360        )361 362        return pn.Row(363            pn.Column(364                controls["x_axis_select"],365                controls["y_axis_select"],366                pn.layout.Divider(margin=(5, 0, 5, 0)),367                controls["color_col_select"],368                controls["color_palette_select"],369                pn.layout.Divider(margin=(5, 0, 5, 0)),370                controls["size_col_select"],371                controls["size_range_slider"],372                controls["size_gamma_slider"],373                pn.layout.Divider(margin=(5, 0, 5, 0)),374                controls["bins_slider"],375                controls["show_gmm"],376                controls["n_components_x"],377                controls["n_components_y"],378                controls["show_linear_fit"],379                pn.layout.Divider(margin=(5, 0, 5, 0)),380                pn.Accordion(381                    (382                        "Plot settings",383                        pn.Column(384                            controls["alpha_slider"],385                            controls["width_slider"],386                            controls["height_slider"],387                            controls["hist_height_slider"],388                            controls["font_size_slider"],389                            width=control_width - 30,390                        ),391                    ),392                    active=[1],393                ),394                margin=(0, 50, 20, 0),  # top, right, bottom, left margins in pixels395                width=control_width,396            ),397            scatter_plot,398            margin=(0, 20, 20, 20),  # top, right, bottom, left margins in pixels399            # width=800,400        )401 402    def _sync_url_state(self, tabs, spike_controls, filter_query=None):403        """Centralize all URL sync logic for the app."""404 405        location = pn.state.location406        location.sync(tabs, {"active": "tab"})407        location.sync(408            self.data_holder,409            {410                "ephys_roi_id_selected": "cell_id",411                "sweep_number_selected": "sweep",412            },413        )414 415        spike_mapping = {416            "extract_from": "spike_extract",417            "spike_type": "spike_type",418            "dim_reduction_method": "dim_method",419            "spike_range": "spike_range",420            "normalize_window_v": "norm_v",421            "normalize_window_dvdt": "norm_dvdt",422            "n_clusters": "n_clusters",423            "if_show_cluster_on_retro": "show_retro",424            "marker_size": "marker_size",425            "alpha_slider": "alpha",426            "plot_width": "plot_width",427            "plot_height": "plot_height",428            "font_size": "font_size",429        }430        for control_name, url_param in spike_mapping.items():431            location.sync(spike_controls[control_name], {"value": url_param})432 433        self.scatter_plot.sync_controls_to_url()434 435        if filter_query is not None:436            location.sync(filter_query, {"value": "query"})437 438    def create_cell_selector_panel(self, filtered_df_meta):439        """440        Builds and returns the cell selector panel that displays metadata.441        """442        # MultiSelect widget to choose additional columns.443        cols = list(filtered_df_meta.columns)444        cols.sort()445        selectable_cols = [col for col in cols if col not in self.cell_key]446        col_selector = pn.widgets.MultiSelect(447            name="Add more columns to show in the table",448            options=selectable_cols,449            value=[450                "ipfx_width_rheo",451                "efel_AP_width @ long_square_rheo, aver",452                "ipfx_sag",453                "efel_sag_ratio1 @ subthreshold, aver",454            ],  # start with no additional columns455            height=300,456            width=500,457        )458 459        def add_df_meta_col(selected_columns, filtered_df_meta):460            """461            Add selected columns from the filtered metadata DataFrame.462            """463            return filtered_df_meta[self.cell_key + selected_columns]464 465        filtered_df_meta = pn.bind(add_df_meta_col, col_selector, filtered_df_meta)466        tab_df_meta = pn.widgets.Tabulator(467            filtered_df_meta,468            selectable=1,469            disabled=True,  # Not editable470            frozen_columns=self.cell_key,471            groupby=["injection region"],472            header_filters=True,473            show_index=False,474            height=300,475            sizing_mode="stretch_width",476            pagination=None,477            stylesheets=[":host .tabulator {font-size: 12px;}"],478        )479 480        # When a row is selected, update the current cell (ephys_roi_id).481        def update_sweep_view_from_table(event):482            if event.new:483                selected_index = event.new[0]484                self.data_holder.ephys_roi_id_selected = str(485                    int(486                        self.data_holder.filtered_df_meta.iloc[selected_index][487                            "ephys_roi_id"488                        ]489                    )490                )491 492        tab_df_meta.param.watch(update_sweep_view_from_table, "selection")493 494        scatter_plot = self.create_scatter_plot()495 496        cell_selector_panel = pn.Column(497            pn.Row(498                col_selector,499                tab_df_meta,500                height=350,501            ),502            pn.Row(503                scatter_plot,504            ),505        )506        return cell_selector_panel507 508    def create_sweep_panel(self, ephys_roi_id=""):509        """510        Builds and returns the sweep visualization panel for a single cell.511        """512        if ephys_roi_id == "":513            return pn.pane.Markdown("Please select a cell from the table above.")514 515        # Load the NWB file for the selected cell.516        raw_this_cell = PatchSeqNWB(ephys_roi_id=ephys_roi_id, if_load_metadata=False)517 518        # Now let's get df sweep from the eFEL enriched one519        df_sweeps = load_efel_features_from_roi(ephys_roi_id, if_from_s3=True)[520            "df_sweeps"521        ]522        df_sweeps_valid = df_sweeps.query("passed == passed")523 524        # Set initial sweep number to first valid sweep525        if self.data_holder.sweep_number_selected == 0:526            self.data_holder.sweep_number_selected = int(527                df_sweeps_valid.iloc[0]["sweep_number"]528            )529 530        # Add a slider to control the downsample factor531        downsample_factor = pn.widgets.IntSlider(532            name="Downsample factor",533            value=5,534            start=1,535            end=10,536        )537 538        # Bind the plotting function to the data holder's sweep number539        bokeh_panel = pn.bind(540            self.update_bokeh,541            raw=raw_this_cell,542            sweep=self.data_holder.param.sweep_number_selected,543            downsample_factor=downsample_factor.param.value_throttled,544        )545 546        # Bind the S3 URL retrieval to the data holder's sweep number547        def get_s3_sweep_images(sweep_number):548            s3_url = get_public_url_sweep(ephys_roi_id, sweep_number)549            images = []550            if isinstance(s3_url, dict) and "sweep" in s3_url:551                images.append(pn.pane.PNG(s3_url["sweep"], width=800, height=400))552            if isinstance(s3_url, dict) and "spikes" in s3_url:553                images.append(pn.pane.PNG(s3_url["spikes"], width=800, height=400))554            return (555                pn.Column(*images)556                if images557                else pn.pane.Markdown("No S3 images available")558            )559 560        s3_sweep_images_panel = pn.bind(561            get_s3_sweep_images,562            sweep_number=self.data_holder.param.sweep_number_selected,563        )564        sweep_pane = pn.Column(565            s3_sweep_images_panel,566            bokeh_panel,567            downsample_factor,568            sizing_mode="stretch_width",569        )570 571        # Build a Tabulator for sweep metadata.572        tab_sweeps = pn.widgets.Tabulator(573            df_sweeps_valid[574                [575                    "sweep_number",576                    "stimulus_code_ext",577                    "stimulus_name",578                    "stimulus_amplitude",579                    "passed",580                    "efel_num_spikes",581                    "num_spikes",582                    "stimulus_start_time",583                    "stimulus_duration",584                    "tags",585                    "reasons",586                    "stimulus_code",587                ]588            ],  # Only show valid sweeps (passed is not NaN)589            hidden_columns=["stimulus_code"],590            selectable=1,591            disabled=True,  # Not editable592            frozen_columns=["sweep_number"],593            header_filters=True,594            show_index=False,595            height=700,596            width=1000,597            groupby=["stimulus_code"],598            stylesheets=[":host .tabulator {font-size: 12px;}"],599        )600 601        # Apply conditional row highlighting.602        if hasattr(tab_sweeps, "style"):603            tab_sweeps.style.apply(604                PatchSeqNWBApp.highlight_selected_rows,605                highlight_subset=df_sweeps_valid.query("passed == True")[606                    "sweep_number"607                ].tolist(),608                color="lightgreen",609                fields=["passed"],610                axis=1,611            ).apply(612                PatchSeqNWBApp.highlight_selected_rows,613                highlight_subset=df_sweeps_valid.query("passed != passed")[614                    "sweep_number"615                ].tolist(),616                color="salmon",617                fields=["passed"],618                axis=1,619            ).apply(620                PatchSeqNWBApp.highlight_selected_rows,621                highlight_subset=df_sweeps_valid.query("passed == False")[622                    "sweep_number"623                ].tolist(),624                color="yellow",625                fields=["passed"],626                axis=1,627            ).apply(628                PatchSeqNWBApp.highlight_selected_rows,629                highlight_subset=df_sweeps_valid.query("num_spikes > 0")[630                    "sweep_number"631                ].tolist(),632                color="lightgreen",633                fields=["num_spikes"],634                axis=1,635            )636 637        # --- Synchronize table selection with sweep number ---638        def update_sweep_from_table(event):639            """Update sweep number when table selection changes."""640            if event.new:641                selected_index = event.new[0]642                new_sweep = int(df_sweeps_valid.iloc[selected_index]["sweep_number"])643                self.data_holder.sweep_number_selected = new_sweep644 645        tab_sweeps.param.watch(update_sweep_from_table, "selection")646        # --- End Synchronization ---647 648        # Build a reactive QC message panel.649        sweep_msg = pn.bind(650            PatchSeqNWBApp.get_qc_message,651            sweep=self.data_holder.param.sweep_number_selected,652            df_sweeps=df_sweeps,653        )654        sweep_msg_panel = pn.pane.Markdown(sweep_msg, width=600, height=30)655 656        return pn.Row(657            pn.Column(658                pn.pane.Markdown(f"# {ephys_roi_id}"),659                pn.pane.Markdown("Select a sweep from the table to view its data."),660                pn.Column(sweep_msg_panel, sweep_pane),661                width=700,662                margin=(0, 100, 0, 0),  # top, right, bottom, left margins663            ),664            pn.Column(665                pn.pane.Markdown("## Sweep metadata"),666                tab_sweeps,667            ),668        )669 670    def main_layout(self):671        """672        Constructs the full application layout with Bootstrap template.673        """674        pn.config.throttled = False675 676        # Create and bind the cell selector panel to filtered metadata677        pane_cell_selector = pn.bind(678            self.create_cell_selector_panel,679            filtered_df_meta=self.data_holder.param.filtered_df_meta,680        )681 682        # Create spike analysis controls and plots683        spike_controls = self.raw_spike_analysis.create_plot_controls()684 685        def update_spike_plots(686            extract_from,687            spike_type,688            n_clusters,689            alpha,690            width,691            height,692            marker_size,693            if_show_cluster_on_retro,694            if_edge_color_projection,695            normalize_window_v,696            normalize_window_dvdt,697            spike_range,698            dim_reduction_method,699            font_size,700            filtered_df_meta=None,701        ):702            df_spikes = self.raw_spike_analysis.get_spikes(spike_type)703            # Extract representative spikes (normalized) - with peak alignment for time-series plot704            df_v_norm, df_dvdt_norm = extract_representative_spikes(705                df_spikes=df_spikes,706                extract_from=extract_from,707                if_normalize_v=True,708                normalize_window_v=normalize_window_v,709                if_normalize_dvdt=True,710                normalize_window_dvdt=normalize_window_dvdt,711                if_smooth_dvdt=False,712                if_align_dvdt_peaks=True,713                filtered_df_meta=filtered_df_meta,714            )715 716            # Extract representative spikes (normalized) - without peak alignment for normalized phase plot717            df_v_norm_phase, df_dvdt_norm_phase = extract_representative_spikes(718                df_spikes=df_spikes,719                extract_from=extract_from,720                if_normalize_v=True,721                normalize_window_v=normalize_window_v,722                if_normalize_dvdt=True,723                normalize_window_dvdt=normalize_window_dvdt,724                if_smooth_dvdt=False,725                if_align_dvdt_peaks=False,726                filtered_df_meta=filtered_df_meta,727            )728 729            # Extract representative spikes (unnormalized) - without peak alignment for phase plots730            df_v_unnorm, df_dvdt_unnorm = extract_representative_spikes(731                df_spikes=df_spikes,732                extract_from=extract_from,733                if_normalize_v=False,734                normalize_window_v=normalize_window_v,735                if_normalize_dvdt=False,736                normalize_window_dvdt=normalize_window_dvdt,737                if_smooth_dvdt=False,738                if_align_dvdt_peaks=False,739                filtered_df_meta=filtered_df_meta,740            )741 742            # Create spike analysis plots with the filtered dataframe743            return self.raw_spike_analysis.create_raw_PCA_plots(744                df_v_norm=df_v_norm,745                df_dvdt_norm=df_dvdt_norm,746                df_v_phase_norm=df_v_norm_phase,747                df_dvdt_phase_norm=df_dvdt_norm_phase,748                df_v_unnorm=df_v_unnorm,749                df_dvdt_unnorm=df_dvdt_unnorm,750                n_clusters=n_clusters,751                alpha=alpha,752                width=width,753                height=height,754                marker_size=marker_size,755                if_show_cluster_on_retro=if_show_cluster_on_retro,756                if_edge_color_projection=if_edge_color_projection,757                spike_range=spike_range,758                dim_reduction_method=dim_reduction_method,759                font_size=font_size,760                normalize_window_v=normalize_window_v,761                normalize_window_dvdt=normalize_window_dvdt,762            )763 764        # Create spike analysis plots765        controls = spike_controls  # shorter name for readability766        param_keys = [767            "n_clusters",768            "alpha_slider",769            "plot_width",770            "if_show_cluster_on_retro",771            "if_edge_color_projection",772            "plot_height",773            "marker_size",774            "normalize_window_v",775            "normalize_window_dvdt",776            "spike_range",777            "dim_reduction_method",778            "font_size",779        ]780        params = {781            k: (782                controls[k].param.value783                if k not in ["if_show_cluster_on_retro", "dim_reduction_method"]784                else controls[k].param.value785            )786            for k in param_keys787        }788 789        spike_plots = pn.bind(790            update_spike_plots,791            extract_from=controls["extract_from"].param.value,792            spike_type=controls["spike_type"].param.value,793            n_clusters=params["n_clusters"],794            alpha=params["alpha_slider"],795            width=params["plot_width"],796            height=params["plot_height"],797            marker_size=params["marker_size"],798            if_show_cluster_on_retro=params["if_show_cluster_on_retro"],799            if_edge_color_projection=params["if_edge_color_projection"],800            normalize_window_v=params["normalize_window_v"],801            normalize_window_dvdt=params["normalize_window_dvdt"],802            spike_range=params["spike_range"],803            dim_reduction_method=params["dim_reduction_method"],804            font_size=params["font_size"],805            filtered_df_meta=self.data_holder.param.filtered_df_meta,806        )807 808        # Create cell summary plot809        def get_s3_cell_summary_plot(ephys_roi_id):810            s3_url = get_public_url_cell_summary(ephys_roi_id)811            if s3_url:812                return pn.pane.PNG(s3_url, sizing_mode="stretch_width")813            else:814                return pn.pane.Markdown(815                    "### Select the table or the scatter plot to view the cell summary plot."816                )817 818        s3_cell_summary_plot = pn.Column(819            pn.bind(820                lambda ephys_roi_id: pn.pane.Markdown(821                    "## Cell summary plot"822                    + (f" for {ephys_roi_id}" if ephys_roi_id else "")823                ),824                ephys_roi_id=self.data_holder.param.ephys_roi_id_selected,825            ),826            pn.bind(827                get_s3_cell_summary_plot,828                ephys_roi_id=self.data_holder.param.ephys_roi_id_selected,829            ),830            sizing_mode="stretch_width",831        )832 833        # Bind the sweep panel to the current cell selection.834        pane_one_cell = pn.bind(835            self.create_sweep_panel,836            ephys_roi_id=self.data_holder.param.ephys_roi_id_selected,837        )838 839        # Create a toggle button for showing/hiding raw sweeps840        show_sweeps_button = pn.widgets.Button(841            name="Show raw sweeps", button_type="primary", width=200842        )843        show_sweeps = pn.widgets.Toggle(name="Show raw sweeps", value=False)844 845        # Link the button to the toggle846        def toggle_sweeps(event):847            show_sweeps.value = not show_sweeps.value848            show_sweeps_button.name = (849                "Hide raw sweeps" if show_sweeps.value else "Show raw sweeps"850            )851 852        show_sweeps_button.on_click(toggle_sweeps)853 854        # Create a dynamic layout that includes pane_one_cell only when show_sweeps is True855        dynamic_content = pn.bind(856            lambda show: pn.Column(pane_one_cell) if show else pn.Column(),857            show_sweeps.param.value,858        )859 860        # --- Connect global filter components ---861        filter_query = pn.widgets.TextAreaInput(862            name="Query string",863            value="`jem-status_reporter` == 'Positive' & `injection region` not in ['Non-Retro', 'Thalamus']",864            placeholder="Enter a pandas query string",865            sizing_mode="stretch_width",866            height=100,867        )868 869        filter_button = pn.widgets.Button(870            name="Apply filter",871            button_type="primary",872            width=150,873        )874 875        reset_button = pn.widgets.Button(876            name="Reset filter",877            button_type="light",878            width=150,879        )880 881        filter_status = pn.pane.Markdown("", css_classes=["alert", "p-2", "m-2"])882 883        # Connect the button to the filter function884        def apply_filter_callback(event):885            result = self.apply_global_filter(filter_query.value)886            if "reset" in result.lower() or "success" in result.lower():887                filter_status.css_classes = ["alert", "alert-success", "p-2", "m-2"]888            elif "error" in result.lower():889                filter_status.css_classes = ["alert", "alert-danger", "p-2", "m-2"]890            else:891                filter_status.css_classes = ["alert", "alert-info", "p-2", "m-2"]892            filter_status.object = result893 894        def reset_filter_callback(event):895            filter_query.value = ""896            result = self.apply_global_filter("")897            filter_status.css_classes = ["alert", "alert-success", "p-2", "m-2"]898            filter_status.object = result899 900        filter_button.on_click(apply_filter_callback)901        reset_button.on_click(reset_filter_callback)902        # --- End filter components ---903 904        # Build the filter panel905        filter_panel = pn.Column(906            pn.pane.Markdown("### Global Filter", css_classes=["card-title"]),907            pn.pane.Markdown(908                """909                    #### Enter a pandas query to filter cells. Examples:910                    LC targeting911                    -  (Only retro cells) `` `injection region` != 'Non-Retro' ``912 913                    Gene QC914                    -  (Nucleus present) `` `jem-nucleus_post_patch` == "nucleus_present" ``915                    -  (QC based on mitochondrial RNA) `` `gene_RNA_QC (log_normed)` == True ``916 917                    Gene expression918                    -  (Fluorescence+) `` `jem-status_reporter` == "Positive" ``919                    -  (Dbh subclass) `` mapmycells_subclass_name.str.contains("DBH", case=False, na=False) ``920                    -  (Marker genes any positive) `` `gene_Dbh (log_normed)` > 0 or `gene_Th (log_normed)` > 0 or921                         `gene_Slc18a2 (log_normed)` > 0 or `gene_Slc6a2 (log_normed)` > 0 ``922                    -  (Dbh+) `` `gene_Dbh (log_normed)` > 0 ``923 924                    Location925                    -  `` `X (A --> P)` > 9500 and `X (A --> P)` < 11500 and926                         `Y (D --> V)` > 2500 and `Y (D --> V)` < 6000 ``927                    """928            ),929            pn.Column(930                filter_query,931                pn.Row(filter_button, reset_button),932            ),933            filter_status,934            width=800,935            margin=(0, 100, 50, 0),  # top, right, bottom, left margins in pixels936            css_classes=["card", "p-4", "m-4"],937        )938 939        if not hasattr(self, "_filter_autoload_registered"):940 941            def _apply_filter_on_load():942                apply_filter_callback(None)943 944            pn.state.onload(_apply_filter_on_load)945            self._filter_autoload_registered = True946 947        # Create the filtered count display948        filtered_count = pn.bind(949            lambda filtered_df: pn.pane.Markdown(950                f"### Filtered cells: {len(filtered_df)} of {len(self.df_meta)} total",951                css_classes=["alert", "alert-info", "p-2", "text-center"],952            ),953            filtered_df=self.data_holder.param.filtered_df_meta,954        )955 956        # Create tabs for different sections957        tabs = pn.Tabs(958            (959                "Cell Explorer",960                pn.Column(961                    filtered_count,962                    pane_cell_selector,963                ),964            ),965            (966                "Spike Analysis",967                pn.Card(968                    pn.Row(969                        pn.Column(970                            pn.pane.Markdown(971                                "### Controls", css_classes=["card-title"]972                            ),973                            *spike_controls.values(),974                            width=250,975                        ),976                        pn.Column(spike_plots),977                    ),978                    title="Raw Spike Analysis",979                    collapsed=False,980                ),981            ),982            (983                "Raw Sweeps",984                pn.Column(985                    show_sweeps_button,986                    pn.pane.Markdown(987                        "### Select a cell from the Cell Explorer tab to view its raw sweeps",988                        css_classes=["alert", "alert-info", "p-2"],989                    ),990                    dynamic_content,991                ),992            ),993            (994                "Feature Distribution",995                pn.Card(996                    pn.pane.PNG(997                        S3_PUBLIC_URL_BASE998                        + "/efel/cell_stats/distribution_all_features.png",999                        width=1300,1000                    ),1001                    title="Distribution of all features",1002                    collapsed=False,1003                ),1004            ),1005            dynamic=True,  # Allow dynamic updates to tab content1006        )1007 1008        self._sync_url_state(tabs, spike_controls, filter_query)1009 1010        # Initialize _active_tab from current tab state (important for URL loading)1011        self._active_tab = tabs.active1012 1013        # Create download button in sidebar1014        self._download_button = pn.widgets.FileDownload(1015            label="Download figures (SVG)",1016            filename="plots.zip",1017            button_type="success",1018            sizing_mode="stretch_width",1019        )1020        self._download_button.callback = self._download_svg_context_aware1021 1022        # Track active tab to change download behavior1023        def update_active_tab(event):1024            self._active_tab = event.new1025 1026        tabs.param.watch(update_active_tab, "active")1027 1028        # Create the template1029        template = pn.template.BootstrapTemplate(1030            title="LC-NE Patch-seq Data Explorer",1031            header_background="#0072B5",  # Allen Institute blue1032            favicon=(1033                "https://alleninstitute.org/wp-content/uploads/2021/10/"1034                "cropped-favicon-32x32.png"1035            ),1036            main=[1037                # pn.pane.Markdown("# Patch-seq Ephys Data Explorer", css_classes=["display-4"]),1038                # pn.layout.Divider(),1039                pn.Row(filter_panel, s3_cell_summary_plot),1040                tabs,1041            ],1042            sidebar=[1043                pn.pane.Markdown("### Filtered Cells"),1044                pn.bind(1045                    lambda filtered_df: pn.pane.Markdown(1046                        f"**{len(filtered_df)} of {len(self.df_meta)} total**",1047                        css_classes=["alert", "alert-info", "p-2"],1048                    ),1049                    filtered_df=self.data_holder.param.filtered_df_meta,1050                ),1051                pn.pane.Markdown("### Selected Cell"),1052                pn.bind(1053                    lambda id: pn.pane.Markdown(1054                        f"**Cell ID:** {id}" if id else "No cell selected",1055                        css_classes=["alert", "alert-secondary", "p-2"],1056                    ),1057                    id=self.data_holder.param.ephys_roi_id_selected,1058                ),1059                pn.bind(1060                    lambda id: pn.pane.Markdown(1061                        f"**Sweep:** {id}" if id else "",1062                        css_classes=["alert", "alert-secondary", "p-2"],1063                    ),1064                    id=self.data_holder.param.sweep_number_selected,1065                ),1066                pn.layout.Divider(margin=(10, 0)),1067                pn.pane.Markdown("### Export"),1068                self._download_button,1069                self._export_progress,1070            ],1071            theme="default",1072        )1073        template.sidebar_width = 2001074        return template1075 1076 1077app = PatchSeqNWBApp()1078layout = app.main_layout()1079layout.servable()1080