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

197 lines
7.6 KiB
Python

#!/usr/bin/env python3
"""
siamese_identify.py — Apply a trained siamese model to IQ-OTH/NCCD to build
a patient manifest via connected-components clustering.
Usage:
conda activate fundus_imaging
python scripts/siamese_identify.py
python scripts/siamese_identify.py --threshold 0.95
python scripts/siamese_identify.py --backbone resnet34
"""
import os
import sys
import csv
import re
import argparse
os.environ["OMP_NUM_THREADS"] = "1"
os.environ["OPENBLAS_NUM_THREADS"] = "1"
os.environ["MKL_NUM_THREADS"] = "1"
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
from classes import SiamesePatientMatcher
# ---------------------------------------------------------------------------
# Config
# ---------------------------------------------------------------------------
ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
MODELS_DIR = os.path.join(ROOT, "models")
RESULTS_DIR = os.path.join(ROOT, "results")
DEFAULT_DATASET = os.path.join(os.path.dirname(ROOT), "The IQ-OTHNCCD lung cancer dataset")
BACKBONE_INPUT_SIZES = {
"resnet18": 224,
"resnet34": 224,
"efficientnet_b0": 240,
}
def f2n(fname):
"""Convert IQ-OTH filename to short form: 'Bengin case (1).jpg''B_001'."""
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
# ---------------------------------------------------------------------------
# Main
# ---------------------------------------------------------------------------
def main():
ap = argparse.ArgumentParser(description=__doc__)
ap.add_argument("--model", default=os.path.join(MODELS_DIR, "siamese_resnet18.pt"),
help="Path to trained siamese model.")
ap.add_argument("--backbone", default="resnet18",
choices=list(BACKBONE_INPUT_SIZES.keys()))
ap.add_argument("--dataset", default=DEFAULT_DATASET,
help="Path to IQ-OTH/NCCD dataset directory.")
ap.add_argument("--threshold", type=float, default=0.9,
help="Minimum siamese probability to create an edge.")
ap.add_argument("--top-k", type=int, default=20,
help="Top-K candidates to verify per slice.")
ap.add_argument("--cluster-method", default="edge_rank",
choices=["edge_rank", "complete", "average"],
help="Clustering on the siamese graph. 'edge_rank' is "
"single-linkage (chains on OOD data); 'complete'/"
"'average' are chaining-resistant agglomerative.")
ap.add_argument("--min-size", type=int, default=None,
help="Absorb groups smaller than this into their nearest "
"group (mitigates leakage-prone singletons).")
ap.add_argument("--max-size", type=int, default=None,
help="Split groups larger than this at their natural gaps "
"in siamese-distance space.")
ap.add_argument("--keep-k", action="store_true",
help="Preserve the known patient count K while enforcing "
"size bounds (balanced bisection + nearest-merge).")
ap.add_argument("--output", default=None,
help="Output manifest path (default: results/siamese_manifest.csv).")
ap.add_argument("--device", default=None)
args = ap.parse_args()
input_size = BACKBONE_INPUT_SIZES[args.backbone]
# ---- Load model ----
print(f"Loading model: {args.model}")
print(f" backbone={args.backbone}, input_size={input_size}")
matcher = SiamesePatientMatcher(
args.model, backbone=args.backbone,
device=args.device, input_size=input_size)
# ---- Collect IQ-OTH images ----
print(f"\nScanning dataset: {args.dataset}")
image_paths = []
class_labels = []
for class_name in sorted(os.listdir(args.dataset)):
class_path = os.path.join(args.dataset, class_name)
if not os.path.isdir(class_path):
continue
for fname in sorted(os.listdir(class_path)):
if fname.lower().endswith((".png", ".jpg", ".jpeg")):
image_paths.append(os.path.join(class_path, fname))
class_labels.append(class_name)
print(f" Found {len(image_paths)} images across "
f"{len(set(class_labels))} classes")
# ---- Identify patients within each class (spectral clustering with known K) ----
KNOWN_K = {
"Bengin cases": 15, # typo in original dataset
"Malignant cases": 40,
"Normal cases": 55,
}
all_assignments = {}
for class_name in sorted(set(class_labels)):
class_mask = [i for i, c in enumerate(class_labels) if c == class_name]
class_paths = [image_paths[i] for i in class_mask]
class_fnames = [os.path.basename(p) for p in class_paths]
k = KNOWN_K.get(class_name)
print(f"\n{'='*50}")
print(f"Class: {class_name} ({len(class_paths)} images, k={k})")
print(f"{'='*50}")
manifest = matcher.identify_patients(
class_paths,
filenames=[f2n(f) for f in class_fnames],
threshold=args.threshold,
top_k=args.top_k,
k=k,
cluster_method=args.cluster_method,
min_size=args.min_size,
max_size=args.max_size,
keep_k=args.keep_k,
)
# Prefix with class and a per-class running index. Enumerate rather than
# reuse the raw cluster id: rebalancing yields pids like "siamese_0_s0"
# whose last token ("s0") is not unique and would collide.
short_cls = {"Bengin cases": "Benign", "Malignant cases": "Malignant",
"Normal cases": "Normal"}[class_name]
prefixed = {f"{short_cls}_{i}": imgs
for i, (pid, imgs) in enumerate(manifest.items())}
all_assignments.update(prefixed)
# ---- Save manifest ----
output_path = args.output or os.path.join(
RESULTS_DIR, "siamese_manifest.csv")
os.makedirs(os.path.dirname(output_path), exist_ok=True)
# Determine class for each patient from the filenames
def get_class(fname):
if fname.startswith("B_"):
return "Benign"
elif fname.startswith("M_"):
return "Malignant"
elif fname.startswith("N_"):
return "Normal"
return "Unknown"
with open(output_path, "w", newline="") as f:
writer = csv.writer(f)
writer.writerow(["patient_id", "class", "n_images", "images"])
for pid in sorted(all_assignments.keys()):
imgs = all_assignments[pid]
cls = get_class(pid)
writer.writerow([pid, cls, len(imgs), ";".join(imgs)])
print(f"\nManifest saved → {output_path}")
print(f" {len(all_assignments)} estimated patients, "
f"{sum(len(v) for v in all_assignments.values())} images")
# Summary per class
print(f"\n{'Class':<20s} {'Patients':>10s} {'Images':>8s} {'Mean imgs/pat':>14s}")
print("-" * 54)
for cls in ["Benign", "Malignant", "Normal"]:
cls_patients = {k: v for k, v in all_assignments.items()
if get_class(k) == cls}
n_pat = len(cls_patients)
n_img = sum(len(v) for v in cls_patients.values())
mean = n_img / n_pat if n_pat > 0 else 0
print(f"{cls:<20s} {n_pat:>10d} {n_img:>8d} {mean:>14.1f}")
print("\nDONE")
if __name__ == "__main__":
main()