"""Plotly visualization functions for cell embedding plots.""" import numpy as np import pandas as pd import plotly.graph_objects as go from .config import ( INFECTION_COLOR_MAP, PLOT_HEIGHT, MARKER_SIZE_NORMAL, MARKER_SIZE_HIGHLIGHTED, MARKER_OPACITY_NORMAL, MARKER_OPACITY_UNKNOWN, ) from . import data def create_embedding_plot( embedding_x="PC1", embedding_y="PC2", highlight_track_id=None, highlight_fov_name=None, ): """Create interactive Plotly scatter plot of embeddings with infection status coloring. Args: embedding_x: X-axis embedding name embedding_y: Y-axis embedding name highlight_track_id: If provided, highlight this specific track highlight_fov_name: If provided with track_id, highlight only this FOV Returns: plotly.graph_objects.Figure: Interactive Plotly figure """ # Get embedding coordinates x_data = data.get_embedding_data(embedding_x) y_data = data.get_embedding_data(embedding_y) # Get metadata for coloring and hover obs_data = data.adata_demo.obs # Map each observation to its color color_values = obs_data["infection_status"].map(INFECTION_COLOR_MAP).values # Map infection status to opacity (unknown cells more transparent) opacity_values = ( obs_data["infection_status"] .map( { "infected": MARKER_OPACITY_NORMAL, "uninfected": MARKER_OPACITY_NORMAL, "unknown": MARKER_OPACITY_UNKNOWN, } ) .values ) # Create DataFrame for plotting plot_data = pd.DataFrame( { "x": x_data, "y": y_data, "color": color_values, "opacity": opacity_values, "id": obs_data["id"].values, "track_id": obs_data["track_id"].values, "timepoint": obs_data["t"].values, "fov_name": obs_data["fov_name"].values, "x_pos": obs_data["x"].values, "y_pos": obs_data["y"].values, "infection_status": obs_data["infection_status"].values, } ) # Store for potential future interactions data.current_plot_data = plot_data # Create Plotly figure fig = go.Figure() # If highlighting a specific track, split into background and foreground if highlight_track_id is not None: # Create filter for the specific track (and FOV if specified) if highlight_fov_name is not None: # Highlight only cells from this specific track in this specific FOV highlight_mask = (plot_data["track_id"] == highlight_track_id) & ( plot_data["fov_name"] == highlight_fov_name ) else: # Highlight all cells with this track_id (across all FOVs) highlight_mask = plot_data["track_id"] == highlight_track_id # Background: all cells NOT in the highlighted set bg_data = plot_data[~highlight_mask] if len(bg_data) > 0: fig.add_trace( go.Scattergl( x=bg_data["x"], y=bg_data["y"], mode="markers", marker=dict( size=3, color="lightgray", opacity=0.2, line=dict(width=0), ), hoverinfo="skip", showlegend=False, ) ) # Foreground: selected track in selected FOV (highlighted) fg_data = plot_data[highlight_mask].copy() if len(fg_data) > 0: # Sort by timepoint to ensure correct trajectory order fg_data = fg_data.sort_values("timepoint").reset_index(drop=True) n_points = len(fg_data) # Add the main trajectory trace with markers and line fig.add_trace( go.Scattergl( x=fg_data["x"], y=fg_data["y"], mode="markers+lines", marker=dict( size=MARKER_SIZE_HIGHLIGHTED, color=fg_data["color"], opacity=0.9, line=dict(width=1, color="white"), ), line=dict(width=2, color="rgba(255, 255, 255, 0.5)"), text=[ f"ID: {row['id']}
" f"Track: {row['track_id']}
" f"Time: {row['timepoint']}
" f"FOV: {row['fov_name']}
" f"Infection: {row['infection_status']}
" f"Position: ({row['x_pos']}, {row['y_pos']})
" f"{embedding_x}: {row['x']:.2f}
" f"{embedding_y}: {row['y']:.2f}" for _, row in fg_data.iterrows() ], hovertemplate="%{text}", showlegend=False, ) ) else: # Standard view: all cells with normal coloring and variable opacity fig.add_trace( go.Scattergl( x=plot_data["x"], y=plot_data["y"], mode="markers", marker=dict( size=MARKER_SIZE_NORMAL, color=plot_data["color"], opacity=plot_data["opacity"], # Use variable opacity per point line=dict(width=0), ), text=[ f"ID: {row['id']}
" f"Track: {row['track_id']}
" f"Time: {row['timepoint']}
" f"FOV: {row['fov_name']}
" f"Infection: {row['infection_status']}
" f"Position: ({row['x_pos']}, {row['y_pos']})
" f"{embedding_x}: {row['x']:.2f}
" f"{embedding_y}: {row['y']:.2f}" for _, row in plot_data.iterrows() ], hovertemplate="%{text}", ) ) # Create title title_text = f"Cell Embedding Visualization: {embedding_x} vs {embedding_y}" if highlight_track_id is not None: if highlight_fov_name is not None: title_text += ( f" (Highlighting Track {highlight_track_id} - {highlight_fov_name})" ) else: title_text += f" (Highlighting Track {highlight_track_id})" fig.update_layout( title=title_text, xaxis_title=embedding_x, yaxis_title=embedding_y, hovermode="closest", height=PLOT_HEIGHT, showlegend=False, template="plotly_dark", ) print(f"Created plot with {len(plot_data)} cells, colored by infection status") return fig