Team Ai
Apppublic

hanhou/patchseq

sourceHugging Faceupdated 8mo agoView on Hugging Face
0likes
spike_analysis.py1626 linesDownload Raw Back to components
1"""2Spike analysis component for the visualization app.3"""4 5from functools import partial6 7import numpy as np8import pandas as pd9import panel as pn10from bokeh.layouts import gridplot, row11from bokeh.models import (12    BoxZoomTool,13    CategoricalColorMapper,14    ColorBar,15    ColumnDataSource,16    CustomJS,17    HoverTool,18    LinearColorMapper,19    Span,20    WheelZoomTool,21)22from bokeh.palettes import Blues256, Inferno256, Reds256, diverging_palette23from bokeh.plotting import figure24from scipy.stats import mannwhitneyu, multivariate_normal, ttest_ind25from sklearn.cluster import KMeans26from sklearn.decomposition import PCA27from sklearn.metrics import silhouette_score28 29try:  # UMAP does not work in Hugging Face Spaces30    from umap import UMAP31except:32    UMAP = None33 34from LCNE_patchseq_analysis import REGION_COLOR_MAPPER35from LCNE_patchseq_analysis.data_util.mesh import trimesh_to_bokeh_data36from LCNE_patchseq_analysis.pipeline_util.s3 import get_public_representative_spikes37from LCNE_patchseq_analysis.pipeline_util.s3 import load_mesh_from_s338 39 40class RawSpikeAnalysis:41    """Handles spike waveform analysis and visualization."""42 43    def __init__(self, df_meta: pd.DataFrame, main_app):44        """Initialize with metadata dataframe."""45        self.main_app = main_app46        self.df_meta = df_meta47        self._latest_figures = {}48        tau_cols = [c for c in df_meta.columns if "ipfx_tau" in c]49        self._tau_col = tau_cols[0] if tau_cols else None50 51        # Load extracted raw spike data52        self.spike_cache = {}53        self.df_spikes = self.get_spikes("average")54        self.extract_from_options = self.df_spikes.index.get_level_values(1).unique()55 56    def get_spikes(self, spike_type: str) -> pd.DataFrame:57        if spike_type not in self.spike_cache:58            self.spike_cache[spike_type] = get_public_representative_spikes(spike_type)59        return self.spike_cache[spike_type]60 61    def create_plot_controls(self) -> dict:62        """Create control widgets for spike analysis."""63        controls = {64            "extract_from": pn.widgets.Select(65                name="Extract spikes from",66                options=sorted(self.extract_from_options.tolist()),67                value="long_square_rheo, min",68                sizing_mode="stretch_width",69            ),70            "spike_type": pn.widgets.Select(71                name="Which spike in a train",72                options=["average", "first", "second", "last"],73                value="average",74                sizing_mode="stretch_width",75            ),76            "dim_reduction_method": pn.widgets.Select(77                name="Dimensionality Reduction Method",78                options=["PCA", "UMAP"],79                value="PCA",80                sizing_mode="stretch_width",81            ),82            "spike_range": pn.widgets.RangeSlider(83                name="Spike Analysis Range (ms)",84                start=-5,85                end=10,86                value=(-3, 6),87                step=0.5,88                sizing_mode="stretch_width",89            ),90            "normalize_window_v": pn.widgets.RangeSlider(91                name="V Normalization Window",92                start=-4,93                end=7,94                value=(-2, 4),95                step=0.5,96                sizing_mode="stretch_width",97            ),98            "normalize_window_dvdt": pn.widgets.RangeSlider(99                name="dV/dt Normalization Window",100                start=-3,101                end=6,102                value=(-2, 0),103                step=0.5,104                sizing_mode="stretch_width",105            ),106            "n_clusters": pn.widgets.IntSlider(107                name="Number of Clusters",108                start=2,109                end=5,110                value=2,111                step=1,112                sizing_mode="stretch_width",113            ),114            "if_show_cluster_on_retro": pn.widgets.Checkbox(115                name="Show type color for Retro",116                value=False,117                sizing_mode="stretch_width",118            ),119            "if_edge_color_projection": pn.widgets.Checkbox(120                name="Edge color by projection target",121                value=True,122                sizing_mode="stretch_width",123            ),124            "marker_size": pn.widgets.IntSlider(125                name="Marker Size",126                start=5,127                end=20,128                value=13,129                step=1,130                sizing_mode="stretch_width",131            ),132            "alpha_slider": pn.widgets.FloatSlider(133                name="Alpha",134                start=0.1,135                end=1.0,136                value=0.3,137                step=0.1,138                sizing_mode="stretch_width",139            ),140            "plot_width": pn.widgets.IntSlider(141                name="Plot Width",142                start=200,143                end=800,144                value=550,145                step=50,146                sizing_mode="stretch_width",147            ),148            "plot_height": pn.widgets.IntSlider(149                name="Plot Height",150                start=200,151                end=800,152                value=550,153                step=50,154                sizing_mode="stretch_width",155            ),156            "font_size": pn.widgets.IntSlider(157                name="Font Size",158                start=8,159                end=24,160                value=10,161                step=1,162                sizing_mode="stretch_width",163            ),164        }165        return controls166 167    def perform_dim_reduction_clustering(168        self, df_v_norm: pd.DataFrame, n_clusters: int = 2, method: str = "PCA"169    ):170        """171        Perform dimensionality reduction and K-means clustering on the voltage traces.172 173        Parameters:174            df_v_norm : pd.DataFrame175                Normalized voltage traces176            n_clusters : int177                Number of clusters for K-means178            method : str179                Dimensionality reduction method ("PCA" or "UMAP")180        """181        v = df_v_norm.values182 183        if method == "PCA":184            # Perform PCA185            reducer = PCA()186            v_proj = reducer.fit_transform(v)187            n_components = 5188            columns = [f"PCA{i}" for i in range(1, n_components + 1)]189        else:  # UMAP190            # Perform UMAP191            reducer = UMAP(n_components=2, random_state=42)192            v_proj = reducer.fit_transform(v)193            n_components = 2194            columns = [f"UMAP{i}" for i in range(1, n_components + 1)]195 196        # K-means clustering197        kmeans = KMeans(n_clusters=n_clusters, random_state=42)198        clusters = kmeans.fit_predict(v_proj[:, :2])199 200        # Calculate metrics201        silhouette_avg = silhouette_score(v_proj[:, :2], clusters)202        metrics = {203            "silhouette_avg": silhouette_avg,204        }205 206        # Save data207        df_v_proj = pd.DataFrame(208            v_proj[:, :n_components], index=df_v_norm.index, columns=columns209        )210 211        # Add cluster information to df_v_norm212        clusters_df = pd.DataFrame(213            clusters, index=df_v_norm.index, columns=["cluster_id"]214        )215        self.df_meta = self.df_meta[216            [col for col in self.df_meta.columns if col != "cluster_id"]217        ].merge(clusters_df, on="ephys_roi_id", how="left")218        df_v_proj = df_v_proj.merge(clusters_df, on="ephys_roi_id", how="left")219        merge_cols = [220            "Date_str",221            "ephys_roi_id",222            "injection region",223            "cell_summary_url",224            "jem-id_cell_specimen",225            "X (A --> P)",226            "Y (D --> V)",227        ]228        # Include ipfx_tau column if available229        tau_cols = [c for c in self.df_meta.columns if "ipfx_tau" in c]230        merge_cols.extend(tau_cols)231        self._tau_col = tau_cols[0] if tau_cols else None232 233        df_v_proj = df_v_proj.merge(234            self.df_meta[merge_cols],235            on="ephys_roi_id",236            how="left",237        )238 239        return df_v_proj, clusters, reducer, metrics240 241    def create_tooltips(242        self,243    ):244        """Create tooltips for the hover tool."""245 246        tau_line = ""247        if self._tau_col:248            tau_field = self._tau_col.replace("{", "{{").replace("}", "}}")249            tau_line = f"""250                    <span style="font-size: 15px;">251                        tau = @{{{tau_field}}}{{0.000}}252                    </span><br>"""253 254        tooltips = f"""255             <div style="text-align: left; flex: auto; white-space: nowrap; margin: 0 10px;256                       border: 2px solid black; padding: 10px;">257                    <span style="font-size: 17px;">258                        <b>@Date_str, @{{injection region}}, @{{ephys_roi_id}},259                            @{{jem-id_cell_specimen}}</b><br>260                    </span>{tau_line}261                    <img src="@cell_summary_url{{safe}}" alt="Cell Summary"262                         style="width: 800px; height: auto;">263             </div>264             """265 266        return tooltips267 268    # Add callback to update ephys_roi_id on point tap269    def update_ephys_roi_id(self, df, attr, old, new):270        if new:271            selected_index = new[0]272            ephys_roi_id = str(int(df["ephys_roi_id"][selected_index]))273            # Update the data holder's ephys_roi_id274            if hasattr(self.main_app, "data_holder"):275                self.main_app.data_holder.ephys_roi_id_selected = ephys_roi_id276 277    def create_raw_PCA_plots(278        self,279        df_v_norm: pd.DataFrame,280        df_dvdt_norm: pd.DataFrame,281        df_v_phase_norm: pd.DataFrame | None = None,282        df_dvdt_phase_norm: pd.DataFrame | None = None,283        df_v_unnorm: pd.DataFrame = None,284        df_dvdt_unnorm: pd.DataFrame = None,285        n_clusters: int = 2,286        alpha: float = 0.3,287        width: int = 400,288        height: int = 400,289        font_size: int = 12,290        marker_size: int = 10,291        if_show_cluster_on_retro: bool = True,292        if_edge_color_projection: bool = True,293        spike_range: tuple = (-4, 7),294        dim_reduction_method: str = "PCA",295        normalize_window_v: tuple = (-2, 4),296        normalize_window_dvdt: tuple = (-2, 0),297    ) -> gridplot:298        """Create plots for spike analysis including dimensionality reduction and clustering."""299        # Filter data based on spike_range300        df_v_norm = df_v_norm.loc[301            :,302            (df_v_norm.columns >= spike_range[0])303            & (df_v_norm.columns <= spike_range[1]),304        ]305        df_dvdt_norm = df_dvdt_norm.loc[306            :,307            (df_dvdt_norm.columns >= spike_range[0])308            & (df_dvdt_norm.columns <= spike_range[1]),309        ]310 311        if df_v_phase_norm is not None:312            df_v_phase_norm = df_v_phase_norm.loc[313                :,314                (df_v_phase_norm.columns >= spike_range[0])315                & (df_v_phase_norm.columns <= spike_range[1]),316            ]317        if df_dvdt_phase_norm is not None:318            df_dvdt_phase_norm = df_dvdt_phase_norm.loc[319                :,320                (df_dvdt_phase_norm.columns >= spike_range[0])321                & (df_dvdt_phase_norm.columns <= spike_range[1]),322            ]323 324        # Filter unnormalized data if provided325        if df_v_unnorm is not None:326            df_v_unnorm = df_v_unnorm.loc[327                :,328                (df_v_unnorm.columns >= spike_range[0])329                & (df_v_unnorm.columns <= spike_range[1]),330            ]331        if df_dvdt_unnorm is not None:332            df_dvdt_unnorm = df_dvdt_unnorm.loc[333                :,334                (df_dvdt_unnorm.columns >= spike_range[0])335                & (df_dvdt_unnorm.columns <= spike_range[1]),336            ]337 338        # Perform dimensionality reduction and clustering339        df_v_proj, clusters, reducer, metrics = self.perform_dim_reduction_clustering(340            df_v_norm, n_clusters, dim_reduction_method341        )342        cluster_colors = ["black", "darkgray", "darkblue", "cyan", "darkorange"][343            :n_clusters344        ]345 346        # Common plot settings347        plot_settings = dict(width=width, height=height)348        legend_groups = {}349 350        def register_renderer(label, renderer):351            if not label or renderer is None:352                return353            legend_groups.setdefault(label, []).append(renderer)354 355        def add_timeseries_mean_sem(fig, df_values, color, label):356            if df_values is None or df_values.empty:357                return358            mean = df_values.mean(axis=0)359            if mean.isna().all():360                return361            sem = df_values.sem(axis=0).fillna(0)362            x_vals = pd.to_numeric(mean.index, errors="coerce")363            valid_mask = ~(np.isnan(x_vals) | np.isnan(mean.values))364            if not valid_mask.any():365                return366            x_vals = x_vals[valid_mask]367            mean_vals = mean.values[valid_mask]368            sem_vals = sem.values[valid_mask]369            legend_label = f"{label} (meanยฑSEM)"370            source = ColumnDataSource(371                {372                    "x": x_vals,373                    "mean": mean_vals,374                    "upper": mean_vals + sem_vals,375                    "lower": mean_vals - sem_vals,376                }377            )378            band = fig.varea(379                x="x",380                y1="lower",381                y2="upper",382                source=source,383                fill_color=color,384                fill_alpha=0.15,385                level="underlay",386            )387            register_renderer(legend_label, band)388            line = fig.line(389                x="x",390                y="mean",391                source=source,392                color=color,393                line_width=3,394                legend_label=legend_label,395            )396            register_renderer(legend_label, line)397 398        def add_phase_mean_sem(399            fig, df_v_values, df_dvdt_values, color, label, n_bins=100400        ):401            if (402                df_v_values is None403                or df_dvdt_values is None404                or df_v_values.empty405                or df_dvdt_values.empty406            ):407                return408            v_vals = df_v_values.to_numpy().astype(float, copy=False).ravel()409            dvdt_vals = df_dvdt_values.to_numpy().astype(float, copy=False).ravel()410            finite_mask = np.isfinite(v_vals) & np.isfinite(dvdt_vals)411            if not finite_mask.any():412                return413            v_vals = v_vals[finite_mask]414            dvdt_vals = dvdt_vals[finite_mask]415            v_min, v_max = np.min(v_vals), np.max(v_vals)416            if v_min == v_max:417                return418            bin_edges = np.linspace(v_min, v_max, n_bins + 1)419            bin_centers = 0.5 * (bin_edges[:-1] + bin_edges[1:])420 421            def render_segment(mask, segment_label, line_dash):422                x_seg = v_vals[mask]423                y_seg = dvdt_vals[mask]424                if x_seg.size < 5:425                    return426                bin_indices = np.digitize(x_seg, bin_edges) - 1427                valid = (bin_indices >= 0) & (bin_indices < n_bins)428                if not valid.any():429                    return430                bin_indices = bin_indices[valid]431                y_seg = y_seg[valid]432                means = []433                sems = []434                centers = []435                for b_idx in range(n_bins):436                    bin_mask = bin_indices == b_idx437                    count = np.count_nonzero(bin_mask)438                    if count < 3:439                        continue440                    values = y_seg[bin_mask]441                    centers.append(bin_centers[b_idx])442                    means.append(np.mean(values))443                    sems.append(444                        np.std(values, ddof=1) / np.sqrt(count) if count > 1 else 0.0445                    )446                if len(centers) < 2:447                    return448                centers = np.array(centers)449                means = np.array(means)450                sems = np.array(sems)451                legend_label = f"{label} (meanยฑSEM)"452                band_source = ColumnDataSource(453                    {454                        "x": np.concatenate([centers, centers[::-1]]),455                        "y": np.concatenate([means + sems, (means - sems)[::-1]]),456                    }457                )458                band = fig.patch(459                    x="x",460                    y="y",461                    source=band_source,462                    fill_color=color,463                    fill_alpha=0.1,464                    line_alpha=0,465                    level="underlay",466                )467                register_renderer(legend_label, band)468                line_source = ColumnDataSource({"x": centers, "y": means})469                line = fig.line(470                    x="x",471                    y="y",472                    source=line_source,473                    color=color,474                    line_width=3,475                    line_dash=line_dash,476                    legend_label=legend_label,477                )478                register_renderer(legend_label, line)479 480            render_segment(dvdt_vals >= 0, "dV/dt > 0", "solid")481            render_segment(dvdt_vals < 0, "dV/dt < 0", "solid")482 483        plots = self._init_spike_subplots(484            dim_reduction_method,485            spike_range,486            normalize_window_v,487            normalize_window_dvdt,488            plot_settings,489        )490        p_embedding = plots["embedding"]491        p_embedding_depth = plots["embedding_depth"]492        p_component_y = plots["component_y"]493        p_tau_xy = plots["tau_xy"]494        p_pc1_projection = plots["pc1_projection"]495        p_tau_projection = plots["tau_projection"]496        p_pc1_histogram = plots["pc1_histogram"]497        p_vm = plots["vm"]498        p_vm_depth = plots["vm_depth"]499        p_dvdt = plots["dvdt"]500        p_dvdt_depth = plots["dvdt_depth"]501        p_phase_norm = plots["phase_norm"]502        p_phase_norm_depth = plots["phase_norm_depth"]503        p_phase = plots["phase"]504 505        self._style_subplots(plots.values(), font_size)506 507        phase_norm_v = df_v_phase_norm if df_v_phase_norm is not None else df_v_norm508        phase_norm_dvdt = (509            df_dvdt_phase_norm if df_dvdt_phase_norm is not None else df_dvdt_norm510        )511 512        # -- Plot PCA scatter with contours --513        # Create a single ColumnDataSource for all clusters514        # If injection region is not "Non-Retro", set color to None515        scatter_renderers = []516 517        for i in df_v_proj["cluster_id"].unique():518            # Add dots519            querystr = "cluster_id == @i"520            group_label = f"Cluster {i + 1}"521            if not if_show_cluster_on_retro:522                querystr += " and `injection region` == 'Non-Retro'"523                group_label += " (Non-Retro)"524 525            group_label += f", n={df_v_proj.query(querystr).shape[0]}"526 527            source = ColumnDataSource(df_v_proj.query(querystr))528            scatter = p_embedding.scatter(529                x=f"{dim_reduction_method}1",530                y=f"{dim_reduction_method}2",531                source=source,532                size=marker_size,533                color=cluster_colors[i],534                alpha=alpha,535                legend_label=group_label,536                hover_color="blue",537                selection_color="blue",538            )539            scatter_renderers.append(scatter)540            register_renderer(group_label, scatter)541 542            # Attach the callback to the selection changes543            source.selected.on_change(544                "indices", partial(self.update_ephys_roi_id, source.data)545            )546 547            # Add contours548            values = (549                df_v_proj.query("cluster_id == @i")550                .loc[:, [f"{dim_reduction_method}1", f"{dim_reduction_method}2"]]551                .values552            )553            mean = np.mean(values, axis=0)554            cov = np.cov(values.T)555            x, y = np.mgrid[556                values[:, 0].min() - 0.5 : values[:, 0].max() + 0.5 : 100j,557                values[:, 1].min() - 0.5 : values[:, 1].max() + 0.5 : 100j,558            ]559            pos = np.dstack((x, y))560            rv = multivariate_normal(mean, cov)561            z = rv.pdf(pos)562            add_counter(563                p_embedding, x, y, z, levels=3, line_color=cluster_colors[i], alpha=1564            )565 566        # Add metrics to the plot567        p_embedding.title.text = (568            f"{dim_reduction_method}\nK-means Clustering (n_clusters = {n_clusters})\n"569            f"Silhouette Score: {metrics['silhouette_avg']:.3f}\n"570        )571        p_embedding.toolbar.active_scroll = p_embedding.select_one(WheelZoomTool)572 573        component_col = f"{dim_reduction_method}1"574        if component_col in df_v_proj.columns:575            self._add_lc_mesh_overlay(p_component_y)576            pc_values = pd.to_numeric(df_v_proj[component_col], errors="coerce")577            if pc_values.notna().any():578                # Optionally map edge color to projection target579                if if_edge_color_projection:580                    regions = df_v_proj["injection region"].unique().tolist()581                    edge_color_mapper = CategoricalColorMapper(582                        factors=regions,583                        palette=[584                            REGION_COLOR_MAPPER.get(r, "black") for r in regions585                        ],586                    )587                    line_color_spec = {588                        "field": "injection region",589                        "transform": edge_color_mapper,590                    }591                    line_width_spec = 1.5592                else:593                    line_color_spec = "black"594                    line_width_spec = 0.5595 596                source = ColumnDataSource(df_v_proj)597                palette = diverging_palette(Blues256, Reds256, 256)598                color_mapper = LinearColorMapper(599                    palette=palette,600                    low=float(pc_values.min()),601                    high=float(pc_values.max()),602                )603                pc_scatter = p_component_y.scatter(604                    x="X (A --> P)",605                    y="Y (D --> V)",606                    source=source,607                    size=marker_size,608                    color={"field": component_col, "transform": color_mapper},609                    line_color=line_color_spec,610                    line_width=line_width_spec,611                    alpha=0.7,612                )613                color_bar = ColorBar(color_mapper=color_mapper, width=8)614                p_component_y.add_layout(color_bar, "right")615                source.selected.on_change(616                    "indices", partial(self.update_ephys_roi_id, source.data)617                )618                hovertool = HoverTool(619                    tooltips=self.create_tooltips(),620                    renderers=[pc_scatter],621                )622                p_component_y.add_tools(hovertool)623            p_component_y.toolbar.active_scroll = p_component_y.select_one(624                WheelZoomTool625            )626            p_component_y.y_range.flipped = True627 628        # --- Tau in X/Y space ---629        tau_cols = [c for c in df_v_proj.columns if "ipfx_tau" in c]630        if tau_cols:631            tau_col = tau_cols[0]632            self._add_lc_mesh_overlay(p_tau_xy)633            tau_values = pd.to_numeric(df_v_proj[tau_col], errors="coerce")634            if tau_values.notna().any():635                df_tau_valid = df_v_proj[tau_values.notna()]636                tau_source = ColumnDataSource(df_tau_valid)637                valid_tau = tau_values.dropna()638                tau_mapper = LinearColorMapper(639                    palette=list(Inferno256),640                    low=float(valid_tau.min()),641                    high=float(valid_tau.max()),642                )643                if if_edge_color_projection:644                    regions = df_v_proj["injection region"].unique().tolist()645                    tau_edge_mapper = CategoricalColorMapper(646                        factors=regions,647                        palette=[648                            REGION_COLOR_MAPPER.get(r, "black") for r in regions649                        ],650                    )651                    tau_line_color = {652                        "field": "injection region",653                        "transform": tau_edge_mapper,654                    }655                    tau_line_width = 1.5656                else:657                    tau_line_color = "black"658                    tau_line_width = 0.5659                tau_scatter = p_tau_xy.scatter(660                    x="X (A --> P)",661                    y="Y (D --> V)",662                    source=tau_source,663                    size=marker_size,664                    color={"field": tau_col, "transform": tau_mapper},665                    line_color=tau_line_color,666                    line_width=tau_line_width,667                    alpha=0.7,668                )669                tau_color_bar = ColorBar(color_mapper=tau_mapper, width=8)670                p_tau_xy.add_layout(tau_color_bar, "right")671                tau_source.selected.on_change(672                    "indices", partial(self.update_ephys_roi_id, tau_source.data)673                )674                tau_hover = HoverTool(675                    tooltips=self.create_tooltips(),676                    renderers=[tau_scatter],677                )678                p_tau_xy.add_tools(tau_hover)679                p_tau_xy.title.text = f"{tau_col} in X/Y space"680            p_tau_xy.toolbar.active_scroll = p_tau_xy.select_one(WheelZoomTool)681            p_tau_xy.y_range.flipped = True682 683        # --- Box plots: PCA1 and tau by projection target ---684        spinal_regions = ["C5", "Spinal cord"]685        cortex_regions = ["Cortex", "PL", "PL, MOs"]686 687        def _fmt_pval(p):688            return f"p={p:.2e}" if p < 0.001 else f"p={p:.3f}"689 690        # PCA1 box plot691        pc1_groups = []692        for grp_label, region_set, color in [693            ("Spinal cord", spinal_regions, REGION_COLOR_MAPPER["Spinal cord"]),694            ("Cortex", cortex_regions, REGION_COLOR_MAPPER["Cortex"]),695        ]:696            mask = df_v_proj["injection region"].isin(region_set)697            sub = df_v_proj.loc[mask]698            vals = pd.to_numeric(sub[component_col], errors="coerce").dropna().values699            if len(vals) > 0:700                pc1_groups.append((grp_label, vals, color))701 702        if len(pc1_groups) == 2:703            self._add_box_strip_plot(704                p_pc1_projection, pc1_groups, marker_size, alpha705            )706            _, pc1_mw = mannwhitneyu(707                pc1_groups[0][1], pc1_groups[1][1], alternative="two-sided"708            )709            _, pc1_tt = ttest_ind(710                pc1_groups[0][1], pc1_groups[1][1], equal_var=False711            )712            p_pc1_projection.title.text = (713                f"{component_col} (on normalized V)\n"714                f"rank-sum {_fmt_pval(pc1_mw)}\nt-test {_fmt_pval(pc1_tt)}"715            )716 717        # Tau box plot718        if self._tau_col and self._tau_col in df_v_proj.columns:719            tau_groups = []720            for grp_label, region_set, color in [721                ("Spinal cord", spinal_regions, REGION_COLOR_MAPPER["Spinal cord"]),722                ("Cortex", cortex_regions, REGION_COLOR_MAPPER["Cortex"]),723            ]:724                mask = df_v_proj["injection region"].isin(region_set)725                sub = df_v_proj.loc[mask]726                vals = pd.to_numeric(sub[self._tau_col], errors="coerce").dropna().values727                if len(vals) > 0:728                    tau_groups.append((grp_label, vals, color))729 730            if len(tau_groups) == 2:731                self._add_box_strip_plot(732                    p_tau_projection, tau_groups, marker_size, alpha733                )734                _, tau_mw = mannwhitneyu(735                    tau_groups[0][1], tau_groups[1][1], alternative="two-sided"736                )737                _, tau_tt = ttest_ind(738                    tau_groups[0][1], tau_groups[1][1], equal_var=False739                )740                p_tau_projection.title.text = (741                    f"ipfx_tau\n"742                    f"rank-sum {_fmt_pval(tau_mw)}\nt-test {_fmt_pval(tau_tt)}"743                )744 745        # Add vertical lines for normalization windows746        p_vm.add_layout(747            Span(748                location=normalize_window_v[0],749                dimension="height",750                line_color="blue",751                line_dash="dashed",752                line_width=2,753            )754        )755        p_vm.add_layout(756            Span(757                location=normalize_window_v[1],758                dimension="height",759                line_color="blue",760                line_dash="dashed",761                line_width=2,762            )763        )764        p_dvdt.add_layout(765            Span(766                location=normalize_window_dvdt[0],767                dimension="height",768                line_color="blue",769                line_dash="dashed",770                line_width=2,771            )772        )773        p_dvdt.add_layout(774            Span(775                location=normalize_window_dvdt[1],776                dimension="height",777                line_color="blue",778                line_dash="dashed",779                line_width=2,780            )781        )782 783        # Add boxzoomtool to Vm and dV/dt plots784        box_zoom_x = BoxZoomTool(dimensions="auto")785        p_vm.add_tools(box_zoom_x)786        p_vm.toolbar.active_drag = box_zoom_x787        box_zoom_x = BoxZoomTool(dimensions="auto")788        p_dvdt.add_tools(box_zoom_x)789        p_dvdt.toolbar.active_drag = box_zoom_x790 791        # Plot voltage and dV/dt traces792        for i in range(n_clusters):793            query_str = "cluster_id == @i"794            group_label = f"Cluster {i + 1}"795            if not if_show_cluster_on_retro:796                query_str += " and `injection region` == 'Non-Retro'"797                group_label += " (Non-Retro)"798            group_label += f", n={df_v_proj.query(query_str).shape[0]}"799            ephys_roi_ids = df_v_proj.query(query_str).ephys_roi_id.tolist()800 801            # Common line properties802            line_props = {803                "alpha": alpha,804                "hover_line_color": "blue",805                "hover_line_alpha": 1.0,806                "hover_line_width": 4,807                "selection_line_color": "blue",808                "selection_line_alpha": 1.0,809                "selection_line_width": 4,810            }811            # Plot voltage traces812            df_this = df_v_norm.query("ephys_roi_id in @ephys_roi_ids")813            source = ColumnDataSource(814                {815                    "xs": [df_v_norm.columns.values] * len(df_this),816                    "ys": df_this.values.tolist(),817                    "ephys_roi_id": ephys_roi_ids,818                }819            )820 821            renderer = p_vm.multi_line(822                source=source,823                xs="xs",824                ys="ys",825                color=cluster_colors[i],826                **line_props,827                legend_label=group_label,828            )829            register_renderer(group_label, renderer)830            add_timeseries_mean_sem(p_vm, df_this, cluster_colors[i], group_label)831 832            # Plot dV/dt traces833            df_this = df_dvdt_norm.query("ephys_roi_id in @ephys_roi_ids")834            source = ColumnDataSource(835                {836                    "xs": [df_dvdt_norm.columns.values] * len(df_this),837                    "ys": df_this.values.tolist(),838                    "ephys_roi_id": ephys_roi_ids,839                }840            )841            renderer = p_dvdt.multi_line(842                source=source,843                xs="xs",844                ys="ys",845                color=cluster_colors[i],846                **line_props,847                legend_label=group_label,848            )849            register_renderer(group_label, renderer)850            add_timeseries_mean_sem(p_dvdt, df_this, cluster_colors[i], group_label)851 852            # Plot phase plot (dV/dt vs V) - normalized853            df_v_this = phase_norm_v.query("ephys_roi_id in @ephys_roi_ids")854            df_dvdt_this = phase_norm_dvdt.query("ephys_roi_id in @ephys_roi_ids")855            source = ColumnDataSource(856                {857                    "xs": df_v_this.values.tolist(),858                    "ys": df_dvdt_this.values.tolist(),859                    "ephys_roi_id": ephys_roi_ids,860                }861            )862            renderer = p_phase_norm.multi_line(863                source=source,864                xs="xs",865                ys="ys",866                color=cluster_colors[i],867                **line_props,868                legend_label=group_label,869            )870            register_renderer(group_label, renderer)871            add_phase_mean_sem(872                p_phase_norm,873                df_v_this,874                df_dvdt_this,875                cluster_colors[i],876                group_label,877            )878 879            # Plot phase plot (dV/dt vs V) - unnormalized880            if df_v_unnorm is not None and df_dvdt_unnorm is not None:881                df_v_unnorm_this = df_v_unnorm.query("ephys_roi_id in @ephys_roi_ids")882                df_dvdt_unnorm_this = df_dvdt_unnorm.query(883                    "ephys_roi_id in @ephys_roi_ids"884                )885                source = ColumnDataSource(886                    {887                        "xs": df_v_unnorm_this.values.tolist(),888                        "ys": df_dvdt_unnorm_this.values.tolist(),889                        "ephys_roi_id": ephys_roi_ids,890                    }891                )892                renderer = p_phase.multi_line(893                    source=source,894                    xs="xs",895                    ys="ys",896                    color=cluster_colors[i],897                    **line_props,898                    legend_label=group_label,899                )900                register_renderer(group_label, renderer)901 902        # Add region cluster_colors to the all plots903        for region in self.df_meta["injection region"].unique():904            if region == "Non-Retro":905                continue906            roi_ids = df_v_proj.query(907                "`injection region` == @region"908            ).ephys_roi_id.tolist()909            legend_label = f"{region}, n={len(roi_ids)}"910 911            source = ColumnDataSource(df_v_proj.query("ephys_roi_id in @roi_ids"))912            scatter = p_embedding.scatter(913                x=f"{dim_reduction_method}1",914                y=f"{dim_reduction_method}2",915                source=source,916                color=REGION_COLOR_MAPPER[region],917                alpha=0.8,918                size=marker_size,919                legend_label=legend_label,920            )921            scatter_renderers.append(scatter)922            register_renderer(legend_label, scatter)923 924            # Attach the callback to the selection changes925            source.selected.on_change(926                "indices", partial(self.update_ephys_roi_id, source.data)927            )928 929            df_v_region = df_v_norm.query("ephys_roi_id in @roi_ids")930            ys = df_v_region.values931 932            # Common line properties933            line_props = {934                "hover_line_color": "blue",935                "hover_line_alpha": 1.0,936                "hover_line_width": 4,937                "selection_line_color": "blue",938                "selection_line_alpha": 1.0,939                "selection_line_width": 4,940            }941            renderer = p_vm.multi_line(942                xs=[df_v_region.columns.values] * ys.shape[0],943                ys=ys.tolist(),944                color=REGION_COLOR_MAPPER[region],945                alpha=0.8,946                legend_label=legend_label,947                **line_props,948            )949            register_renderer(legend_label, renderer)950            add_timeseries_mean_sem(951                p_vm, df_v_region, REGION_COLOR_MAPPER[region], legend_label952            )953 954            df_dvdt_region = df_dvdt_norm.query("ephys_roi_id in @roi_ids")955            ys = df_dvdt_region.values956            renderer = p_dvdt.multi_line(957                xs=[df_dvdt_region.columns.values] * ys.shape[0],958                ys=ys.tolist(),959                color=REGION_COLOR_MAPPER[region],960                alpha=0.8,961                legend_label=legend_label,962                **line_props,963            )964            register_renderer(legend_label, renderer)965            add_timeseries_mean_sem(966                p_dvdt, df_dvdt_region, REGION_COLOR_MAPPER[region], legend_label967            )968 969            # Plot phase plot (dV/dt vs V) for regions - normalized970            df_v_norm_region = phase_norm_v.query("ephys_roi_id in @roi_ids")971            df_dvdt_norm_region = phase_norm_dvdt.query("ephys_roi_id in @roi_ids")972            v_vals_norm = df_v_norm_region.values973            dvdt_vals_norm = df_dvdt_norm_region.values974            renderer = p_phase_norm.multi_line(975                xs=v_vals_norm.tolist(),976                ys=dvdt_vals_norm.tolist(),977                color=REGION_COLOR_MAPPER[region],978                alpha=0.8,979                legend_label=legend_label,980                **line_props,981            )982            register_renderer(legend_label, renderer)983            add_phase_mean_sem(984                p_phase_norm,985                df_v_norm_region,986                df_dvdt_norm_region,987                REGION_COLOR_MAPPER[region],988                legend_label,989            )990 991            # Plot phase plot (dV/dt vs V) for regions - unnormalized992            if df_v_unnorm is not None and df_dvdt_unnorm is not None:993                v_vals_unnorm = df_v_unnorm.query("ephys_roi_id in @roi_ids").values994                dvdt_vals_unnorm = df_dvdt_unnorm.query(995                    "ephys_roi_id in @roi_ids"996                ).values997                renderer = p_phase.multi_line(998                    xs=v_vals_unnorm.tolist(),999                    ys=dvdt_vals_unnorm.tolist(),1000                    color=REGION_COLOR_MAPPER[region],1001                    alpha=0.8,1002                    legend_label=legend_label,1003                    **line_props,1004                )1005                register_renderer(legend_label, renderer)1006 1007        depth_values = pd.to_numeric(df_v_proj["Y (D --> V)"], errors="coerce")1008        if depth_values.notna().any():1009            depth_mapper = LinearColorMapper(1010                palette=list(reversed(Inferno256)),1011                low=float(depth_values.min()),1012                high=float(depth_values.max()),1013            )1014            valid_mask = depth_values.notna()1015            depth_source = ColumnDataSource(df_v_proj.loc[valid_mask])1016            depth_scatter = p_embedding_depth.scatter(1017                x=f"{dim_reduction_method}1",1018                y=f"{dim_reduction_method}2",1019                source=depth_source,1020                size=marker_size,1021                color={"field": "Y (D --> V)", "transform": depth_mapper},1022                line_color="black",1023                line_width=0.3,1024                alpha=0.8,1025            )1026            color_bar = ColorBar(color_mapper=depth_mapper, width=8)1027            p_embedding_depth.add_layout(color_bar, "right")1028            depth_source.selected.on_change(1029                "indices", partial(self.update_ephys_roi_id, depth_source.data)1030            )1031            depth_hover = HoverTool(1032                tooltips=self.create_tooltips(),1033                renderers=[depth_scatter],1034            )1035            p_embedding_depth.add_tools(depth_hover)1036 1037            missing_mask = ~valid_mask1038            if missing_mask.any():1039                missing_source = ColumnDataSource(df_v_proj.loc[missing_mask])1040                missing_scatter = p_embedding_depth.scatter(1041                    x=f"{dim_reduction_method}1",1042                    y=f"{dim_reduction_method}2",1043                    source=missing_source,1044                    size=marker_size,1045                    color="gray",1046                    line_color="black",1047                    line_width=0.3,1048                    alpha=0.7,1049                    legend_label="Depth missing",1050                )1051                missing_source.selected.on_change(1052                    "indices", partial(self.update_ephys_roi_id, missing_source.data)1053                )1054                missing_hover = HoverTool(1055                    tooltips=self.create_tooltips(),1056                    renderers=[missing_scatter],1057                )1058                p_embedding_depth.add_tools(missing_hover)1059 1060            depth_map = df_v_proj.set_index("ephys_roi_id")["Y (D --> V)"]1061            roi_ids = df_v_norm.index.tolist()1062            depth_series = depth_map.reindex(roi_ids)1063 1064            if not depth_series.isna().all():1065                depth_line_props = {1066                    "hover_line_color": "blue",1067                    "hover_line_alpha": 1.0,1068                    "hover_line_width": 4,1069                    "selection_line_color": "blue",1070                    "selection_line_alpha": 1.0,1071                    "selection_line_width": 4,1072                }1073                missing_ids = depth_series[depth_series.isna()].index.tolist()1074                valid_ids = depth_series[depth_series.notna()].index.tolist()1075                vm_source = ColumnDataSource(1076                    {1077                        "xs": [df_v_norm.columns.values] * len(valid_ids),1078                        "ys": df_v_norm.loc[valid_ids].values.tolist(),1079                        "depth": depth_series.loc[valid_ids].tolist(),1080                        "ephys_roi_id": valid_ids,1081                    }1082                )1083                p_vm_depth.multi_line(1084                    source=vm_source,1085                    xs="xs",1086                    ys="ys",1087                    line_color={"field": "depth", "transform": depth_mapper},1088                    alpha=0.8,1089                    **depth_line_props,1090                )1091                if missing_ids:1092                    df_v_missing = df_v_norm.loc[missing_ids]1093                    p_vm_depth.multi_line(1094                        xs=[df_v_missing.columns.values] * len(df_v_missing),1095                        ys=df_v_missing.values.tolist(),1096                        line_color="gray",1097                        alpha=0.6,1098                        legend_label="Depth missing",1099                        **depth_line_props,1100                    )1101                dvdt_source = ColumnDataSource(1102                    {1103                        "xs": [df_dvdt_norm.columns.values] * len(valid_ids),1104                        "ys": df_dvdt_norm.loc[valid_ids].values.tolist(),1105                        "depth": depth_series.loc[valid_ids].tolist(),1106                        "ephys_roi_id": valid_ids,1107                    }1108                )1109                p_dvdt_depth.multi_line(1110                    source=dvdt_source,1111                    xs="xs",1112                    ys="ys",1113                    line_color={"field": "depth", "transform": depth_mapper},1114                    alpha=0.8,1115                    **depth_line_props,1116                )1117                if missing_ids:1118                    df_dvdt_missing = df_dvdt_norm.loc[missing_ids]1119                    p_dvdt_depth.multi_line(1120                        xs=[df_dvdt_missing.columns.values] * len(df_dvdt_missing),1121                        ys=df_dvdt_missing.values.tolist(),1122                        line_color="gray",1123                        alpha=0.6,1124                        legend_label="Depth missing",1125                        **depth_line_props,1126                    )1127 1128                phase_source = ColumnDataSource(1129                    {1130                        "xs": phase_norm_v.reindex(valid_ids).values.tolist(),1131                        "ys": phase_norm_dvdt.reindex(valid_ids).values.tolist(),1132                        "depth": depth_series.loc[valid_ids].tolist(),1133                        "ephys_roi_id": valid_ids,1134                    }1135                )1136                p_phase_norm_depth.multi_line(1137                    source=phase_source,1138                    xs="xs",1139                    ys="ys",1140                    line_color={"field": "depth", "transform": depth_mapper},1141                    alpha=0.8,1142                    **depth_line_props,1143                )1144                if missing_ids:1145                    df_v_phase_missing = phase_norm_v.loc[missing_ids]1146                    df_dvdt_phase_missing = phase_norm_dvdt.loc[missing_ids]1147                    p_phase_norm_depth.multi_line(1148                        xs=df_v_phase_missing.values.tolist(),1149                        ys=df_dvdt_phase_missing.values.tolist(),1150                        line_color="gray",1151                        alpha=0.6,1152                        legend_label="Depth missing",1153                        **depth_line_props,1154                    )1155 1156        # Add tooltips1157        # Add renderers like this to solve bug like this:1158        #   File "/Users/han.hou/miniconda3/envs/patch-seq/lib/python3.10/1159        # site-packages/panel/io/location.py", line 57, in _get_location_params1160        #     params['pathname'], search = uri.split('?')1161        # ValueError: too many values to unpack (expected 2)1162        # 2025-04-09 00:03:04,658 500 GET /patchseq_panel_viz??? (::1) 8541.01ms1163        hovertool = HoverTool(1164            tooltips=self.create_tooltips(),1165            renderers=scatter_renderers,1166        )1167        p_embedding.add_tools(hovertool)1168 1169        hovertool = HoverTool(1170            tooltips=[("ephys_roi_id", "@ephys_roi_id")],1171            attachment="right",  # Fix tooltip to the right of the plot1172        )1173        p_vm.add_tools(hovertool)1174        p_dvdt.add_tools(hovertool)1175        p_vm_depth.add_tools(hovertool)1176        p_dvdt_depth.add_tools(hovertool)1177 1178        hovertool = HoverTool(1179            tooltips=[("ephys_roi_id", "@ephys_roi_id")],1180            attachment="right",1181        )1182        p_phase_norm.add_tools(hovertool)1183        p_phase_norm_depth.add_tools(hovertool)1184 1185        hovertool = HoverTool(1186            tooltips=[("ephys_roi_id", "@ephys_roi_id")],1187            attachment="right",1188        )1189        p_phase.add_tools(hovertool)1190 1191        # Add boxzoomtool to phase plot1192        box_zoom_x = BoxZoomTool(dimensions="auto")1193        p_phase.add_tools(box_zoom_x)1194        p_phase.toolbar.active_drag = box_zoom_x1195 1196        box_zoom_x = BoxZoomTool(dimensions="auto")1197        p_phase_norm.add_tools(box_zoom_x)1198        p_phase_norm.toolbar.active_drag = box_zoom_x1199 1200        legend_configs = {

Showing the first 1,200 of 1626 lines. Download the file for the rest.