hanhou/patchseq
0
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 = {