"""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