AURAD / detection /plot_boxplot.py
diing's picture
Upload folder using huggingface_hub
41c8683 verified
Raw
History Blame Contribute Delete
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()