CoolFace
Apppublic

Marcel0123/unsupervised-training

sourceHugging Facemitupdated 1y agoView on Hugging Face
0likes
app.py321 linesDownload Raw Back to root
1 2import gradio as gr3import numpy as np4import pandas as pd5import matplotlib.pyplot as plt6from sklearn import datasets7from sklearn.preprocessing import StandardScaler8from sklearn.decomposition import PCA9import plotly.graph_objects as go10import plotly.express as px11import time12 13FEATURE_LABELS = {14    "age": "Leeftijd",15    "sex": "Geslacht",16    "bmi": "BMI (Body Mass Index)",17    "bp": "Bloeddruk",18    "s1": "Totale cholesterol",19    "s2": "LDL-cholesterol",20    "s3": "HDL-cholesterol",21    "s4": "Chol./HDL-verhouding",22    "s5": "Triglyceriden",23    "s6": "Bloedsuiker (glucose)",24    "target": "Doelscore (progressie)",25}26LABEL_TO_KEY = {v: k for k, v in FEATURE_LABELS.items()}27 28MEDICAL_MD = """29### Medisch nut  30 31**Wat zien we hier?**  32Ik heb een bestaande, anonieme gezondheidsdataset gebruikt die speciaal beschikbaar is gemaakt voor onderzoek en studie. In deze gegevens staan metingen van een grote groep patiënten, zoals **bloedwaarden, BMI, cholesterol en bloedsuiker**.  33 34Zo'n enorme berg cijfers is voor artsen en ziekenhuizen bijna niet in één keer te overzien. Het is gewoon te veel om met het blote oog patronen uit te halen.  35 36**Daar komt kunstmatige intelligentie om de hoek kijken.**  37Met deze techniek (PCA) kan de computer de data slim samenvatten en patronen zichtbaar maken. Dit programma dat ik heb ontworpen laat live zien hoe die samenvatting werkt.  38 39- Elke punt is één patiënt.  40- De kleur laat zien hoe hoog of laag een bepaalde meting is (standaard: BMI).  41- De pijlen (in de 2D-biplot) laten zien welke metingen het meeste invloed hebben.  42- Links bovenin kun je kiezen welke meting je als uitgangspunt wilt nemen.  43 44**En wat heb je hieraan?**  45In de praktijk gebruiken artsen en onderzoekers zo'n plot om patronen en verbanden te ontdekken. 👉 Het is dus niet alleen een mooi plaatje, maar echt een manier om grote hoeveelheden data sneller en slimmer te begrijpen.  46 47Met AI kunnen we patronen vinden die je met het blote oog nooit zou zien. Dat maakt dit niet alleen een mooie visualisatie, maar ook een knap stukje technologie met échte waarde voor onderzoek en zorg.  48 49**Speel zelf de onderzoeker!**  50Doe alsof je een arts bent en kies links bovenin een waarde, bijvoorbeeld **cholesterol**, **leeftijd** of **geslacht**. Klik daarna op **Update visualisaties** en ontdek je eigen patronen in de data.  51"""52 53# -------------------- Data helpers --------------------54def load_diabetes_df():55    d = datasets.load_diabetes()56    X = pd.DataFrame(d.data, columns=d.feature_names)  # gestandaardiseerd57    y = pd.Series(d.target, name="target")58    df = X.copy(); df["target"] = y59    return df60 61def compute_overview_table(df: pd.DataFrame):62    keys = ["bmi","bp","s1","s2","s3","s4","s5","s6"]63    rows = []64    for k in keys:65        vals = df[k].dropna().values66        mean = float(vals.mean())67        pct_above = float((vals > 0).mean() * 100.0)  # 0 ≈ globaal gemiddelde68        pct_below = float((vals < 0).mean() * 100.0)69        rows.append({70            "Meting": FEATURE_LABELS.get(k, k),71            "Gemiddelde (gestandaardiseerd)": round(mean, 3),72            "% boven gemiddelde": round(pct_above, 1),73            "% onder gemiddelde": round(pct_below, 1),74        })75    table = pd.DataFrame(rows)76    note = ("Let op: waarden in deze dataset zijn **gestandaardiseerd**. `0` ≈ algemeen gemiddelde. "77            "Positief = hoger dan gemiddeld, negatief = lager dan gemiddeld.")78    return table, note79 80# -------------------- PCA helpers --------------------81def compute_pca(df: pd.DataFrame, n_components: int, standardize: bool):82    feats = [c for c in df.columns if c != "target"]83    X = df[feats].values84    if standardize:85        scaler = StandardScaler(with_mean=True, with_std=True)86        Xs = scaler.fit_transform(X)87    else:88        Xs = X89    pca = PCA(n_components=min(int(n_components), Xs.shape[1]))90    Z = pca.fit_transform(Xs)91    loadings = pca.components_.T92    expl = pca.explained_variance_ratio_93    return feats, Xs, Z, loadings, expl94 95# -------------------- Plot builders --------------------96def build_biplot_plotly(df, Z, loadings, feats, color_key, arrow_scale=2.0):97    # Hover info98    fields = ["bmi","bp","s1","s2","s3","s4","s5","s6","age","sex","target"]99    hover_text = [100        "<br>".join(f"{FEATURE_LABELS.get(k,k)}: {df.iloc[i][k]:.3f}" for k in fields)101        for i in range(len(df))102    ]103    fig = go.Figure()104    fig.add_trace(go.Scatter(105        x=Z[:,0], y=Z[:,1], mode="markers",106        marker=dict(size=8, color=df[color_key].values),107        text=hover_text, hovertemplate="%{text}<extra></extra>"108    ))109    # loading pijlen110    for i, key in enumerate(feats):111        x = loadings[i,0]*arrow_scale; y = loadings[i,1]*arrow_scale112        fig.add_annotation(x=x, y=y, ax=0, ay=0, xref="x", yref="y", axref="x", ayref="y",113                           showarrow=True, arrowhead=3)114        fig.add_annotation(x=x*1.05, y=y*1.05, text=FEATURE_LABELS.get(key,key),115                           showarrow=False, font=dict(size=10))116    fig.update_layout(title="PCA-biplot (2D, hover)", xaxis_title="PC1", yaxis_title="PC2",117                      margin=dict(l=10, r=10, t=40, b=10))118    return fig119 120def build_biplot_matplotlib(df, Z, loadings, feats, color_key, arrow_scale=2.0, point_size=32, alpha=0.85):121    fig = plt.figure()122    ax = fig.add_subplot(111)123    sc = ax.scatter(Z[:,0], Z[:,1], c=df[color_key].values, s=point_size, alpha=alpha)124    cbar = plt.colorbar(sc, ax=ax, pad=0.02); cbar.set_label(f"Kleur: {FEATURE_LABELS.get(color_key,color_key)}")125    ax.set_xlabel("PC1"); ax.set_ylabel("PC2"); ax.set_title("PCA-biplot — PNG-export")126    for i,key in enumerate(feats):127        x=loadings[i,0]*arrow_scale; y=loadings[i,1]*arrow_scale128        ax.arrow(0,0,x,y, head_width=0.05, head_length=0.08, fc="k", ec="k", length_includes_head=True)129        ax.text(x*1.08, y*1.08, FEATURE_LABELS.get(key,key), fontsize=9, ha="center", va="center")130    ax.axhline(0,color="grey",linewidth=0.6,linestyle=":"); ax.axvline(0,color="grey",linewidth=0.6,linestyle=":")131    ax.grid(True,linestyle=":",linewidth=0.6); fig.tight_layout()132    return fig133 134def build_pca3d(Z3, color_vals):135    fig = go.Figure(data=[go.Scatter3d(x=Z3[:,0], y=Z3[:,1], z=Z3[:,2], mode="markers",136                                       marker=dict(size=4, color=color_vals, opacity=0.85))])137    fig.update_layout(title="PCA 3D — PC1·PC2·PC3 (sleep om te draaien)",138                      scene=dict(xaxis_title="PC1", yaxis_title="PC2", zaxis_title="PC3"),139                      margin=dict(l=10, r=10, t=40, b=10))140    return fig141 142def build_variance_plot(expl):143    fig = plt.figure()144    ax = fig.add_subplot(111)145    xs = np.arange(1, len(expl)+1)146    ax.bar(xs, expl, width=0.8, align="center")147    ax.plot(xs, np.cumsum(expl), marker="o")148    ax.set_xticks(xs); ax.set_xlabel("Principal Component"); ax.set_ylabel("Explained variance ratio")149    ax.set_title("Uitlegvariantie per component (balken) + cumulatief (lijn)")150    ax.grid(True, linestyle=":", linewidth=0.6); fig.tight_layout()151    return fig152 153def build_hist_box(df: pd.DataFrame, color_key: str):154    series = df[color_key].dropna()155    label = FEATURE_LABELS.get(color_key, color_key)156    fig_hist = px.histogram(x=series, nbins=30, title=f"Histogram — {label}", labels={"x": label})157    fig_hist.update_layout(xaxis_title=label, yaxis_title="Aantal", margin=dict(l=10, r=10, t=40, b=10))158    fig_box = px.box(y=series, points="outliers", title=f"Boxplot — {label}", labels={"y": label})159    fig_box.update_layout(yaxis_title=label, margin=dict(l=10, r=10, t=40, b=10))160    return fig_hist, fig_box161 162# -------------------- Controllers --------------------163def controller(color_label="BMI (Body Mass Index)", n_components=10, standardize=True, arrow_scale=2.0):164    df = load_diabetes_df()165    feats, Xs, Z, loadings, expl = compute_pca(df, n_components, standardize)166    color_key = LABEL_TO_KEY.get(color_label, "bmi")167    color_vals = df[color_key].values168 169    fig_biplot = build_biplot_plotly(df, Z, loadings, feats, color_key, arrow_scale=arrow_scale)170    if Z.shape[1] < 3:171        pca3 = PCA(n_components=3); Z3 = pca3.fit_transform(Xs)172    else:173        Z3 = Z[:, :3]174    fig3d = build_pca3d(Z3, color_vals)175    fig_variance = build_variance_plot(expl)176    fig_hist, fig_box = build_hist_box(df, color_key)177 178    load_df = pd.DataFrame({179        "feature_key": feats,180        "PC1_loading": loadings[:, 0],181        "PC2_loading": loadings[:, 1],182        "PC1_abs": np.abs(loadings[:, 0]),183        "PC2_abs": np.abs(loadings[:, 1]),184    })185    load_df["Feature (PC1)"] = load_df["feature_key"].map(lambda k: FEATURE_LABELS.get(k, k))186    load_df["Feature (PC2)"] = load_df["feature_key"].map(lambda k: FEATURE_LABELS.get(k, k))187    top_pc1 = load_df.sort_values("PC1_abs", ascending=False)[["Feature (PC1)", "PC1_loading"]].head(6).reset_index(drop=True)188    top_pc2 = load_df.sort_values("PC2_abs", ascending=False)[["Feature (PC2)", "PC2_loading"]].head(6).reset_index(drop=True)189    max_len = max(len(top_pc1), len(top_pc2))190    top_pc1 = top_pc1.reindex(range(max_len)); top_pc2 = top_pc2.reindex(range(max_len))191    table = pd.concat([top_pc1, top_pc2], axis=1)192 193    overview_df, overview_note = compute_overview_table(df)194 195    summary_md = f"""196### Wat zie je hier?197- **Klik op _Update visualisaties_** om alles te verversen met jouw keuze.198- **Hover** over punten voor exacte waarden (BMI, bloeddruk, cholesterol, glucose, leeftijd, geslacht, etc.).199- **2D-biplot** met pijlen (belangrijkste metingen) en **3D-view** voor extra diepte.200- **Uitlegvariantieplot**: laat zien hoeveel variatie elke component uitlegt.201- **Histogram + boxplot**: verdeling en spreiding van de gekozen meting ({FEATURE_LABELS.get(color_key,color_key)}).202"""203    return fig_biplot, fig3d, fig_variance, table, overview_df, overview_note, summary_md204 205def animate_pca(color_label="BMI (Body Mass Index)", point_size=32, alpha=0.85, n_components=10, standardize=True, frames=40, pause=0.0):206    df = load_diabetes_df()207    feats, Xs, Z, loadings, expl = compute_pca(df, n_components, standardize)208    color_key = LABEL_TO_KEY.get(color_label, "bmi")209    color_vals = df[color_key].values210    for i in range(frames):211        t = i / max(1, frames-1)212        w1 = min(1.0, t * 2.0); w2 = max(0.0, (t - 0.5) * 2.0)213        coords = np.column_stack([Z[:, 0] * w1, Z[:, 1] * w2])214        fig = plt.figure()215        ax = fig.add_subplot(111)216        ax.scatter(coords[:, 0], coords[:, 1], c=color_vals, s=point_size, alpha=alpha)217        ax.set_xlabel("PC1 (opbouw)"); ax.set_ylabel("PC2 (opbouw)")218        title = "PCA-projectie (animatie) — " + ("PC1 →" if w2 == 0 else "PC1 + PC2")219        ax.set_title(f"{title} — frame {i+1}/{frames}")220        ax.axhline(0, color="grey", linewidth=0.6, linestyle=":"); ax.axvline(0, color="grey", linewidth=0.6, linestyle=":")221        ax.grid(True, linestyle=":", linewidth=0.6); fig.tight_layout()222        yield fig223        if pause > 0:224            time.sleep(pause)225 226def export_biplot_png(color_label="BMI (Body Mass Index)", arrow_scale=2.0, point_size=32, alpha=0.85, n_components=10, standardize=True):227    df = load_diabetes_df()228    feats, Xs, Z, loadings, expl = compute_pca(df, n_components, standardize)229    color_key = LABEL_TO_KEY.get(color_label, "bmi")230    fig = build_biplot_matplotlib(df, Z, loadings, feats, color_key, arrow_scale=arrow_scale, point_size=point_size, alpha=alpha)231    path = f"/mnt/data/biplot_{int(time.time())}.png"232    fig.savefig(path, dpi=150, bbox_inches="tight"); plt.close(fig)233    return path234 235def export_variance_png(n_components=10, standardize=True):236    df = load_diabetes_df()237    feats, Xs, Z, loadings, expl = compute_pca(df, n_components, standardize)238    fig = build_variance_plot(expl)239    path = f"/mnt/data/variance_{int(time.time())}.png"240    fig.savefig(path, dpi=150, bbox_inches="tight"); plt.close(fig)241    return path242 243# -------------------- UI --------------------244with gr.Blocks(title="PCA Dashboard — Diabetes (netjes & compleet)") as demo:245    gr.HTML("""246    <style>247      .callout {padding:12px 14px; border-left:4px solid #2563eb; background:#f1f5f9; border-radius:8px; margin: 8px 0 18px;}248      .cta {padding:10px 12px; border:1px dashed #2563eb; background:#eff6ff; border-radius:8px; margin-top:6px;}249    </style>250    """)251 252    gr.Markdown("# PCA Dashboard — Diabetes (netjes & compleet)")253    gr.Markdown(MEDICAL_MD)254    gr.HTML('<div class="callout"><b>Belangrijk:</b> kies links je instellingen en klik daarna op <b>Update visualisaties</b>.     Wil je de stap-voor-stap projectie zien? Klik op <b>▶ Animate PCA</b>.</div>')255 256    with gr.Row():257        with gr.Column(scale=1):258            with gr.Group():259                gr.Markdown("### Instellingen")260                color_choices = [FEATURE_LABELS[k] for k in ["bmi","bp","s1","s2","s3","s4","s5","s6","age","sex","target"]]261                color_feat = gr.Dropdown(choices=color_choices, value=FEATURE_LABELS["bmi"], label="Kleur op meting")262                n_components = gr.Slider(3, 10, value=10, step=1, label="Aantal PCA-componenten")263                standardize = gr.Checkbox(value=True, label="Standaardiseer metingen (aanbevolen)")264                arrow_scale = gr.Slider(0.5, 5.0, value=2.0, step=0.1, label="Pijl-schaal (2D-biplot)")265                run_btn = gr.Button("🔄 Update visualisaties")266                gr.HTML('<div class="cta"><b>Klik hierna op: "🔄 Update visualisaties"</b> om alle grafieken te verversen.</div>')267            with gr.Group():268                gr.Markdown("### Animatie")269                animate_btn = gr.Button("▶ Animate PCA (PC1 → PC2)")270                gr.HTML('<div class="cta"><b>Klik op: "▶ Animate PCA"</b> om de projectie stap-voor-stap te zien.</div>')271                anim_plot = gr.Plot(label="Animatie van projectie")272            with gr.Group():273                gr.Markdown("### Downloads")274                dl_biplot = gr.DownloadButton("Download biplot (PNG)")275                dl_var = gr.DownloadButton("Download variatieplot (PNG)")276 277        with gr.Column(scale=2):278            with gr.Row():279                with gr.Column():280                    gr.Markdown("### Biplot (2D, hover)")281                    plot_biplot = gr.Plotly()282                with gr.Column():283                    gr.Markdown("### 3D PCA (PC1–PC3)")284                    plot3d = gr.Plotly()285            with gr.Row():286                with gr.Column():287                    gr.Markdown("### Uitlegvariantie")288                    plot_expl = gr.Plot()289                with gr.Column():290                    gr.Markdown("### Top-features (PC1 / PC2)")291                    table = gr.Dataframe(headers=["Feature (PC1)", "Loading PC1", "Feature (PC2)", "Loading PC2"], row_count=6)292            with gr.Row():293                with gr.Column():294                    gr.Markdown("### Histogram")295                    plot_hist = gr.Plotly()296                with gr.Column():297                    gr.Markdown("### Boxplot")298                    plot_box = gr.Plotly()299            with gr.Row():300                with gr.Column():301                    gr.Markdown("### Overzicht (gemiddelden & verdeling)")302                    overview_tbl = gr.Dataframe(interactive=False)303                with gr.Column():304                    gr.Markdown("### Samenvatting")305                    summary = gr.Markdown()306                    overview_note_md = gr.Markdown()307 308    inputs = [color_feat, n_components, standardize, arrow_scale]309    run_btn.click(fn=controller, inputs=inputs,310                  outputs=[plot_biplot, plot3d, plot_expl, table, overview_tbl, overview_note_md, summary])311    demo.load(fn=controller, inputs=inputs,312              outputs=[plot_biplot, plot3d, plot_expl, table, overview_tbl, overview_note_md, summary])313 314    animate_btn.click(fn=animate_pca, inputs=[color_feat], outputs=anim_plot)315 316    dl_biplot.click(fn=export_biplot_png, inputs=[color_feat, arrow_scale], outputs=[dl_biplot])317    dl_var.click(fn=export_variance_png, inputs=[], outputs=[dl_var])318 319if __name__ == "__main__":320    demo.queue().launch(server_name="0.0.0.0", server_port=7860, ssr_mode=False, show_api=False)321