Files
patient_leakage_detection/scripts/visualizations/siamese_similarity_tsne.py
T
rpotter6298 35cbd9ac3c 2026001
2026-07-01 17:35:58 +02:00

137 lines
5.5 KiB
Python

#!/usr/bin/env python3
"""siamese_similarity_tsne.py — t-SNE of the siamese's OWN similarity space.
Unlike manifest_tsne.py (which recolors the VGG16 feature layout), this builds
the layout directly from the siamese pairwise distance 1 - P(same-patient), so
proximity reflects how the siamese model itself relates images. Points are
colored by the siamese edge-rank groups.
This is the diagnostic view for the over-merge: if the siamese collapses several
patients together (over-confidence on IQ-OTH), the largest group forms one dense
mass in its own space; coherent patients form tight, separated islands.
Usage:
conda activate fundus_imaging
python scripts/visualizations/siamese_similarity_tsne.py
"""
import os
import sys
import re
import csv
import argparse
import numpy as np
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt
from sklearn.manifold import TSNE
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.dirname(
os.path.abspath(__file__)))))
ROOT = os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
PLOTS_DIR = os.path.join(ROOT, "plots")
DEFAULT_DATASET = os.path.join(os.path.dirname(ROOT),
"The IQ-OTHNCCD lung cancer dataset")
from classes import SiamesePatientMatcher
CLASS_DIRS = {"Bengin cases": "Benign", "Malignant cases": "Malignant",
"Normal cases": "Normal"}
VALID_EXT = (".png", ".jpg", ".jpeg", ".tif", ".tiff", ".bmp")
RANDOM_STATE = 42
def f2n(fname):
m = re.search(r"\((\d+)\)", fname)
num = int(m.group(1)) if m else None
for cls_key, prefix in [("Bengin cases", "B"), ("Malignant cases", "M"),
("Normal cases", "N")]:
if fname.startswith(cls_key.rstrip("s")):
return f"{prefix}_{num:03d}" if num else fname
return fname
def load_manifest(path):
mapping = {}
with open(path, newline="") as f:
for row in csv.DictReader(f):
for img in row["images"].split(";"):
if img:
mapping[img] = row["patient_id"]
return mapping
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--model", default=os.path.join(ROOT, "models", "siamese_resnet18.pt"))
ap.add_argument("--backbone", default="resnet18")
ap.add_argument("--dataset", default=DEFAULT_DATASET)
ap.add_argument("--manifest", default=os.path.join(ROOT, "results", "siamese_manifest.csv"))
ap.add_argument("--name", default="siamese_sim")
ap.add_argument("--highlight-largest", type=int, default=0,
help="Grey all points and bold only the N largest groups.")
ap.add_argument("--tag", default="")
args = ap.parse_args()
tag = f"_{args.tag}" if args.tag else ""
matcher = SiamesePatientMatcher(args.model, backbone=args.backbone, input_size=224)
groups = load_manifest(args.manifest)
for class_dir, cls in CLASS_DIRS.items():
cpath = os.path.join(args.dataset, class_dir)
files = sorted(f for f in os.listdir(cpath) if f.lower().endswith(VALID_EXT))
paths = [os.path.join(cpath, f) for f in files]
ids = [f2n(f) for f in files]
print(f"\n{cls}: {len(paths)} images")
# Siamese pairwise distance -> t-SNE on precomputed distances.
emb = matcher.embed_images(paths)
P = matcher._dense_prob_matrix(emb)
dist = np.clip(1.0 - P, 0.0, None)
np.fill_diagonal(dist, 0.0)
perp = max(5, min(30, (len(paths) - 1) // 3))
X = TSNE(n_components=2, metric="precomputed", init="random",
perplexity=perp, random_state=RANDOM_STATE).fit_transform(dist)
img_group = np.array([groups.get(i, "unassigned") for i in ids])
g_order = sorted(set(img_group), key=lambda g: -(img_group == g).sum())
n = len(g_order)
fig, ax = plt.subplots(figsize=(14, 10))
if args.highlight_largest > 0:
# Grey everything, then bold only the N largest groups.
ax.scatter(X[:, 0], X[:, 1], c="lightgray", s=12, alpha=0.5)
hl = g_order[:args.highlight_largest]
hl_cmap = plt.cm.tab10
for gi, g in enumerate(hl):
gm = img_group == g
ax.scatter(X[gm, 0], X[gm, 1], c=[hl_cmap(gi)], s=28, alpha=0.9,
edgecolors="black", linewidths=0.3,
label=f"{g} ({gm.sum()} imgs)")
ax.legend(loc="lower right", title="Largest siamese groups")
ax.set_title(f"{cls} — siamese similarity t-SNE "
f"(largest {len(hl)} of {n} groups highlighted)")
else:
cmap = plt.cm.tab20 if n <= 20 else plt.cm.gist_ncar
for gi, g in enumerate(g_order):
color = cmap(gi % 20) if n <= 20 else cmap(gi / max(n - 1, 1))
gm = img_group == g
ax.scatter(X[gm, 0], X[gm, 1], c=[color], s=18, alpha=0.8)
biggest = (img_group == g_order[0]).sum()
ax.set_title(f"{cls} — siamese similarity t-SNE "
f"({n} groups; largest={biggest} imgs)")
ax.set_xlabel("t-SNE dim 1 (siamese distance)")
ax.set_ylabel("t-SNE dim 2 (siamese distance)")
plt.tight_layout()
suffix = "_highlight" if args.highlight_largest > 0 else ""
out = os.path.join(PLOTS_DIR, "tsne", f"tsne_{args.name}_{cls}{suffix}{tag}.png")
os.makedirs(os.path.dirname(out), exist_ok=True)
plt.savefig(out, dpi=150)
plt.close()
print(f" Saved {out}")
print("DONE")
if __name__ == "__main__":
main()