#!/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()