Team Ai
Apppublic

Dany546/interactive_plotting

sourceHugging Faceupdated 10mo agoView on Hugging Face
0likes
app.py123 linesDownload Raw Back to root
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