File size: 9,496 Bytes
41c8683 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 | """
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() |