| """ |
| 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: |
| |
| 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() |
|
|
| |
| 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.") |
|
|
| |
| df_disease = df[df["disease"].isin(args.diseases)].copy() |
| if df_disease.empty: |
| raise SystemExit("No rows match the selected --diseases.") |
|
|
| |
| |
| |
| 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) |
|
|
| |
| |
| 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] |
|
|
| |
| 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: |
| |
| ax.set_ylim(-0.02, 1.08) |
|
|
| |
| 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, |
| ) |
|
|
| |
| if args.ref_model is not None and args.ref_model in args.models: |
| n_models = len(args.models) |
| width = 0.75 |
| |
| def box_x(gi, mi): |
| return gi - width / 2 + (mi + 0.5) * width / n_models |
|
|
| ref_idx = args.models.index(args.ref_model) |
|
|
| |
| |
| 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 |
| row_gap = 0.075 |
| 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) |
|
|
| |
| 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 = (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() |