File size: 4,553 Bytes
1c223a0 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 | from __future__ import annotations
import pandas as pd
import plotly.express as px
from rag_chunk_visualizer.core.logging import get_logger
logger = get_logger(__name__)
class PlotBuildError(RuntimeError):
pass
def build_plot_dataframe(projection_df: pd.DataFrame) -> pd.DataFrame:
if projection_df.empty:
raise PlotBuildError("Projection dataframe is empty.")
plot_df = projection_df.copy()
# Keep a small, fixed point size so one document doesn't visually swallow others.
plot_df["point_size"] = 10.0
# Add tiny deterministic jitter only for exact duplicate coordinates.
# This makes overlapping points visible without changing the underlying structure much.
duplicate_rank = plot_df.groupby(["x", "y"]).cumcount()
plot_df["x_plot"] = plot_df["x"] + (duplicate_rank * 0.01)
plot_df["y_plot"] = plot_df["y"] + (duplicate_rank * 0.01)
return plot_df
def create_embedding_scatter_plot(
projection_df: pd.DataFrame,
projection_summary: dict | None = None,
selected_chunk_id: str | None = None,
retrieved_chunk_ids: list[str] | None = None,
query_projection: dict | None = None,
query_text: str | None = None,
):
plot_df = build_plot_dataframe(projection_df)
retrieved_chunk_ids = retrieved_chunk_ids or []
method_label = "2D Projection"
if projection_summary is not None:
method_label = str(projection_summary.get("method", "2D Projection")).upper()
fig = px.scatter(
plot_df,
x="x_plot",
y="y_plot",
color="filename",
symbol="filename",
size="point_size",
hover_name="chunk_id",
hover_data={
"filename": True,
"chunk_index": True,
"char_count": True,
"word_count": True,
"preview": True,
"point_size": False,
"x_plot": False,
"y_plot": False,
"x": ":.4f",
"y": ":.4f",
},
custom_data=["chunk_id", "doc_id", "plot_index"],
render_mode="svg",
title=f"{method_label} embedding map",
)
fig.update_traces(
marker={
"size": 10,
"line": {"width": 1},
"opacity": 0.6,
},
)
if retrieved_chunk_ids:
retrieved_df = plot_df[plot_df["chunk_id"].isin(retrieved_chunk_ids)]
if not retrieved_df.empty:
fig.add_scatter(
x=retrieved_df["x_plot"],
y=retrieved_df["y_plot"],
mode="markers",
name="Retrieved",
customdata=retrieved_df[["chunk_id", "doc_id", "plot_index"]].values,
marker={
"symbol": "circle-open",
"size": 22,
"line": {"width": 3, "color": "#F59E0B"},
},
hovertemplate="Retrieved: %{customdata[0]}<extra></extra>",
)
if selected_chunk_id:
selected_df = plot_df[plot_df["chunk_id"] == selected_chunk_id]
if not selected_df.empty:
fig.add_scatter(
x=selected_df["x_plot"],
y=selected_df["y_plot"],
mode="markers",
name="Selected",
customdata=selected_df[["chunk_id", "doc_id", "plot_index"]].values,
marker={
"symbol": "diamond-open",
"size": 24,
"line": {"width": 3, "color": "#FFFFFF"},
},
hovertemplate="Selected: %{customdata[0]}<extra></extra>",
)
if query_projection is not None:
query_label = query_text.strip() if isinstance(query_text, str) else "Query"
fig.add_scatter(
x=[query_projection["x"]],
y=[query_projection["y"]],
mode="markers+text",
text=["Query"],
textposition="top center",
name="Query",
customdata=[["__query__", "__query__", -1]],
marker={
"symbol": "x",
"size": 18,
"color": "#22C55E",
"line": {"width": 2},
},
hovertemplate=f"Query: {query_label}<extra></extra>",
)
fig.update_layout(
height=560,
margin={"l": 20, "r": 20, "t": 50, "b": 20},
xaxis_title="Projection X",
yaxis_title="Projection Y",
legend_title_text="Legend",
dragmode="select",
clickmode="event+select",
)
return fig
|