Download detection/plot_boxplot.py from diing/AURAD: direct link, hf CLI and curl.
- Browser
- Download file 9.5 kB
-
https://huggingface.co/diing/AURAD/resolve/main/detection/plot_boxplot.py
- Command line
-
hf download hf://diing/AURAD/detection/plot_boxplot.py
-
curl -L -o plot_boxplot.py https://huggingface.co/diing/AURAD/resolve/main/detection/plot_boxplot.py
9.5 kB
| """ | |
| BiomedParse-style grouped boxplot. | |
| X-axis: "All" + selected diseases, each labeled with sample count (n=...) | |
| NOTE: "All" aggregates ONLY the diseases passed via --diseases, | |
| not every disease present in the CSV. | |
| Y-axis: the chosen metric (Dice / Mask_IoU / Box_IoU) | |
| Colors: one color per model, boxes dodged within each group | |
| Top: horizontal legend | |
| Optional significance bracket per group: best-model vs each-other-model, | |
| Wilcoxon signed-rank if paired (same sample_id), else Mann-Whitney U. | |
| Usage: | |
| python plot_boxplot_biomedparse.py \ | |
| --csv per_sample_results.csv \ | |
| --metric Dice \ | |
| --diseases Nodule Mass Effusion Pneumothorax Cardiomegaly \ | |
| --models Baseline DCN Aug Ours Ours+ \ | |
| --ref_model "Ours+" \ | |
| --outdir figs/ | |
| """ | |
| import argparse | |
| import os | |
| from itertools import combinations | |
| import matplotlib.pyplot as plt | |
| import numpy as np | |
| import pandas as pd | |
| import seaborn as sns | |
| from scipy import stats | |
| def parse_args(): | |
| p = argparse.ArgumentParser() | |
| p.add_argument("--csv", required=True) | |
| p.add_argument("--metric", default="Dice", | |
| choices=["Dice", "Mask_IoU", "Box_IoU", "Detected@0.5"]) | |
| p.add_argument("--diseases", nargs="+", required=True, | |
| help="Diseases to show (in display order). 'All' is auto-added at front " | |
| "and aggregates ONLY these selected diseases.") | |
| p.add_argument("--models", nargs="*", default=None, | |
| help="Model display order. Defaults to order in CSV.") | |
| p.add_argument("--ref_model", default=None, | |
| help="If set, run significance tests comparing this model to each other " | |
| "model within each disease group. Stars drawn over the box.") | |
| p.add_argument("--paired", action="store_true", | |
| help="Use Wilcoxon signed-rank (paired by sample_id) instead of Mann-Whitney U.") | |
| p.add_argument("--no_all", action="store_true", | |
| help="Don't prepend the 'All' aggregate column.") | |
| p.add_argument("--ylim", nargs=2, type=float, default=None, | |
| help="Y-axis limits, e.g. --ylim 0 1.05") | |
| p.add_argument("--outdir", default="figs") | |
| p.add_argument("--figsize", nargs=2, type=float, default=None) | |
| return p.parse_args() | |
| def stars(p): | |
| if p < 1e-4: return "****" | |
| if p < 1e-3: return "***" | |
| if p < 1e-2: return "**" | |
| if p < 5e-2: return "*" | |
| return "ns" | |
| def sig_test(a, b, paired): | |
| """Return (p_value, n_used). Handles paired/unpaired and edge cases.""" | |
| a = np.asarray(a, dtype=float) | |
| b = np.asarray(b, dtype=float) | |
| if paired: | |
| # Align by length; assumes caller already aligned by sample_id | |
| m = min(len(a), len(b)) | |
| a, b = a[:m], b[:m] | |
| d = a - b | |
| d = d[~np.isnan(d)] | |
| if len(d) < 3 or np.all(d == 0): | |
| return 1.0, len(d) | |
| try: | |
| stat, p = stats.wilcoxon(d, zero_method="wilcox", alternative="two-sided") | |
| except ValueError: | |
| return 1.0, len(d) | |
| return float(p), len(d) | |
| else: | |
| a = a[~np.isnan(a)] | |
| b = b[~np.isnan(b)] | |
| if len(a) < 3 or len(b) < 3: | |
| return 1.0, min(len(a), len(b)) | |
| stat, p = stats.mannwhitneyu(a, b, alternative="two-sided") | |
| return float(p), min(len(a), len(b)) | |
| def main(): | |
| args = parse_args() | |
| os.makedirs(args.outdir, exist_ok=True) | |
| df = pd.read_csv(args.csv) | |
| df = df[df["metric"] == args.metric].copy() | |
| # Restrict to chosen models (and pin their order) | |
| if args.models is None: | |
| args.models = list(dict.fromkeys(df["model"].tolist())) | |
| df = df[df["model"].isin(args.models)] | |
| if df.empty: | |
| raise SystemExit("No rows after filtering.") | |
| # Restrict to chosen diseases for the per-disease columns | |
| df_disease = df[df["disease"].isin(args.diseases)].copy() | |
| if df_disease.empty: | |
| raise SystemExit("No rows match the selected --diseases.") | |
| # Build the "All" aggregate by relabeling disease -> "All" | |
| # IMPORTANT: aggregate over the SELECTED diseases only (df_disease), | |
| # not the full df, so "All" reflects the diseases shown on the plot. | |
| if not args.no_all: | |
| df_all = df_disease.copy() | |
| df_all["disease"] = "All" | |
| df_plot = pd.concat([df_all, df_disease], ignore_index=True) | |
| group_order = ["All"] + list(args.diseases) | |
| else: | |
| df_plot = df_disease | |
| group_order = list(args.diseases) | |
| # n for x-tick labels (count of unique samples per group, across all models) | |
| # Use the first model to count samples per group (they should all see the same test set) | |
| ref_for_n = args.models[0] | |
| n_per_group = (df_plot[df_plot["model"] == ref_for_n] | |
| .groupby("disease")["sample_id"].nunique().to_dict()) | |
| xticklabels = [f"{g}\n(n = {n_per_group.get(g, 0):,})" for g in group_order] | |
| # --- Plot --- | |
| sns.set_theme(style="whitegrid", context="talk") | |
| palette = sns.color_palette("Set2", n_colors=len(args.models)) | |
| figsize = args.figsize or (max(11, 1.6 * len(group_order) + 4), 6.5) | |
| fig, ax = plt.subplots(figsize=figsize) | |
| sns.boxplot( | |
| data=df_plot, x="disease", y="value", hue="model", | |
| order=group_order, hue_order=args.models, | |
| palette=palette, | |
| showfliers=True, | |
| fliersize=2.5, | |
| linewidth=1.0, | |
| width=0.75, | |
| ax=ax, | |
| ) | |
| ax.set_xticks(range(len(group_order))) | |
| ax.set_xticklabels(xticklabels, rotation=25, ha="right") | |
| ax.set_xlabel("") | |
| ax.set_ylabel(f"{args.metric} score" if args.metric != "Detected@0.5" | |
| else args.metric) | |
| if args.ylim: | |
| ax.set_ylim(*args.ylim) | |
| else: | |
| # Nice default for Dice/IoU | |
| ax.set_ylim(-0.02, 1.08) | |
| # Top horizontal legend (like the BiomedParse figure) | |
| handles, labels = ax.get_legend_handles_labels() | |
| ax.legend( | |
| handles, labels, | |
| loc="lower center", bbox_to_anchor=(0.5, 1.02), | |
| ncol=min(len(args.models), 4), | |
| frameon=False, handlelength=1.5, columnspacing=1.5, | |
| fontsize=11, title=None, | |
| ) | |
| # --- Significance: ref_model vs each other model, per group --- | |
| if args.ref_model is not None and args.ref_model in args.models: | |
| n_models = len(args.models) | |
| width = 0.75 | |
| # x position of a (group_idx, model_idx) box | |
| def box_x(gi, mi): | |
| return gi - width / 2 + (mi + 0.5) * width / n_models | |
| ref_idx = args.models.index(args.ref_model) | |
| # Per-group base y is the local maximum of that group's data | |
| # (lets stars sit just above each group rather than at a global top) | |
| local_top = (df_plot.groupby("disease")["value"].max() | |
| .reindex(group_order).to_dict()) | |
| bump_per_group = {g: 0 for g in group_order} | |
| bracket_h = 0.018 # short tick height | |
| row_gap = 0.075 # vertical gap between stacked brackets (in axes data units) | |
| top_needed = 0.0 | |
| for gi, g in enumerate(group_order): | |
| sub = df_plot[df_plot["disease"] == g] | |
| ref_vals = (sub[sub["model"] == args.ref_model] | |
| .sort_values("sample_id")["value"].values) | |
| base_y = (local_top.get(g, 0.9) or 0.9) + 0.04 | |
| for mi, m in enumerate(args.models): | |
| if m == args.ref_model: | |
| continue | |
| other_vals = (sub[sub["model"] == m] | |
| .sort_values("sample_id")["value"].values) | |
| if len(ref_vals) == 0 or len(other_vals) == 0: | |
| continue | |
| p, _ = sig_test(ref_vals, other_vals, paired=args.paired) | |
| s = stars(p) | |
| if s == "ns": | |
| continue | |
| x1 = box_x(gi, ref_idx) | |
| x2 = box_x(gi, mi) | |
| y = base_y + row_gap * bump_per_group[g] | |
| bump_per_group[g] += 1 | |
| top_needed = max(top_needed, y + bracket_h + 0.03) | |
| ax.plot([x1, x1, x2, x2], | |
| [y, y + bracket_h, y + bracket_h, y], | |
| lw=1.0, color="black") | |
| ax.text((x1 + x2) / 2, y + bracket_h + 0.005, s, | |
| ha="center", va="bottom", fontsize=10) | |
| # Expand ylim to fit stars (but cap at a sane upper bound for Dice/IoU) | |
| if not args.ylim: | |
| ax.set_ylim(-0.02, max(1.08, top_needed)) | |
| sns.despine() | |
| plt.tight_layout() | |
| out = os.path.join(args.outdir, f"boxplot_{args.metric}_biomedparse.png") | |
| plt.savefig(out, dpi=220, bbox_inches="tight") | |
| plt.savefig(out.replace(".png", ".pdf"), bbox_inches="tight") | |
| print(f"Saved {out} (+ .pdf)") | |
| # Summary table | |
| summary = (df_plot.groupby(["disease", "model"])["value"] | |
| .agg(["count", "mean", "median", "std"]).round(4) | |
| .reset_index()) | |
| summary["disease"] = pd.Categorical(summary["disease"], | |
| categories=group_order, ordered=True) | |
| summary["model"] = pd.Categorical(summary["model"], | |
| categories=args.models, ordered=True) | |
| summary = summary.sort_values(["disease", "model"]) | |
| summary_path = os.path.join(args.outdir, f"summary_{args.metric}.csv") | |
| summary.to_csv(summary_path, index=False) | |
| print(f"Saved summary -> {summary_path}") | |
| if __name__ == "__main__": | |
| main() |