""" 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()