Dany546/interactive_plotting
0
1import dash2from dash import dcc, html, Input, Output, State3import plotly.graph_objects as go4import plotly.express as px5import pandas as pd6import numpy as np7import sqlite38import wandb9import os10 11app = dash.Dash(__name__)12 13app.layout = html.Div([14 html.H3("Embedding Explorer (Plotly + Dash + wandb)"),15 html.Div([16 dcc.Input(id="run-path", type="text", placeholder="entity/project/run_id", style={"width": "40%"}),17 dcc.Input(id="artifact-name", type="text", placeholder="embeddings:latest", style={"width": "30%"}),18 html.Button("Download SQL", id="btn-download")19 ], style={"marginBottom": "12px"}),20 21 html.Div(id="status", style={"marginBottom": "12px", "color": "#555"}),22 23 dcc.Checklist(24 id="visible-cats",25 options=[],26 value=[],27 inline=True28 ),29 30 dcc.Graph(id="embedding-plot", style={"height": "800px"}),31 dcc.Store(id="df-store")32])33 34@app.callback(35 Output("df-store", "data"),36 Output("status", "children"),37 Input("btn-download", "n_clicks"),38 State("run-path", "value"),39 State("artifact-name", "value")40)41def download_and_load(n, run_path, artifact_name):42 if not n:43 return dash.no_update, "Enter run path and artifact name."44 try:45 api = wandb.Api()46 run = api.run(run_path)47 artifact = run.use_artifact(artifact_name)48 art_dir = artifact.download()49 db_path = [os.path.join(art_dir, f) for f in os.listdir(art_dir) if f.endswith(".db")][0]50 conn = sqlite3.connect(db_path)51 df = pd.read_sql("SELECT * FROM embeddings", conn)52 conn.close()53 return df.to_json(date_format="iso", orient="split"), f"Loaded {len(df)} rows."54 except Exception as e:55 return dash.no_update, f"Error: {e}"56 57@app.callback(58 Output("embedding-plot", "figure"),59 Output("visible-cats", "options"),60 Output("visible-cats", "value"),61 Input("df-store", "data"),62 Input("visible-cats", "value")63)64def update_plot(df_json, visible_cats):65 if not df_json:66 return go.Figure(), [], []67 df = pd.read_json(df_json, orient="split")68 69 # Identify category columns (multi-column counts)70 cat_cols = [c for c in df.columns if c.startswith("cat_")]71 if not visible_cats:72 visible_cats = cat_cols73 74 counts_matrix = df[cat_cols].values75 emb = df[["tsne_x", "tsne_y"]].values76 77 # Hover text: top 4 counts78 hover_texts = []79 for row in counts_matrix:80 top_idx = np.argsort(row)[::-1][:4]81 top_info = [f"{cat_cols[i]}: {row[i]}" for i in top_idx if row[i] > 0]82 hover_texts.append("<br>".join(top_info))83 84 # Multi-category mode85 if len(visible_cats) > 1:86 visible_idx = [cat_cols.index(c) for c in visible_cats]87 sub_counts = counts_matrix[:, visible_idx]88 max_idx = sub_counts.argmax(axis=1)89 max_cat = [visible_cats[i] for i in max_idx]90 color_map = {c: px.colors.qualitative.Set2[i % len(px.colors.qualitative.Set2)] for i, c in enumerate(cat_cols)}91 point_colors = [color_map[c] for c in max_cat]92 93 fig = go.Figure(go.Scattergl(94 x=emb[:,0], y=emb[:,1],95 mode="markers",96 marker=dict(size=8, color=point_colors, opacity=0.9),97 text=hover_texts,98 hovertemplate="%{text}<extra></extra>"99 ))100 101 # Single-category mode102 elif len(visible_cats) == 1:103 cat = visible_cats[0]104 counts = df[cat].values105 sizes = 6 + 2*np.log1p(counts)106 opacities = 0.3 + 0.6*(counts > 0)107 fig = go.Figure(go.Scattergl(108 x=emb[:,0], y=emb[:,1],109 mode="markers",110 marker=dict(size=sizes, color="steelblue", opacity=opacities),111 text=hover_texts,112 hovertemplate="%{text}<extra></extra>"113 ))114 else:115 fig = go.Figure()116 117 fig.update_layout(height=700, plot_bgcolor="white", title="Embedding Visualization")118 return fig, [{"label": c, "value": c} for c in cat_cols], visible_cats119 120if __name__ == "__main__":121 # Hugging Face Spaces expects port 7860122 app.run_server(host="0.0.0.0", port=7860, debug=False)123 