Download scripts/plot_anat_tsne.py from AIGeeksGroup/HMCLIP: direct link, hf CLI and curl.
- Browser
- Download file 28.8 kB
-
https://huggingface.co/AIGeeksGroup/HMCLIP/resolve/main/scripts/plot_anat_tsne.py
- Command line
-
hf download hf://AIGeeksGroup/HMCLIP/scripts/plot_anat_tsne.py
-
curl -L -o plot_anat_tsne.py https://huggingface.co/AIGeeksGroup/HMCLIP/resolve/main/scripts/plot_anat_tsne.py
28.8 kB
| #!/usr/bin/env python3 | |
| """ | |
| Disease-colored t-SNE for RadIR/RadIR image+text embeddings. | |
| Always jointly embeds Image and Text on ONE figure. | |
| Color = disease / class | |
| Marker = ○ Image, △ Text (same disease uses related hues) | |
| Default --label_mode exclusive: | |
| Keep samples with exactly one label (Normal or a single disease), | |
| drop multi-disease / unlabeled; plot all Image+Text together. | |
| Usage: | |
| python scripts/plot_anat_tsne.py \\ | |
| --load_feats outputs/tsne_feats.npz \\ | |
| --label_mode exclusive \\ | |
| --out outputs/tsne_disease_all.png | |
| """ | |
| from __future__ import annotations | |
| import argparse | |
| import json | |
| import os | |
| import sys | |
| from collections import Counter | |
| import matplotlib.pyplot as plt | |
| import numpy as np | |
| import pandas as pd | |
| import torch | |
| from sklearn.manifold import TSNE | |
| from torch.utils.data import DataLoader | |
| from transformers import AutoModel, AutoTokenizer | |
| ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) | |
| RAD_IR_ROOT = os.path.join(ROOT, "Rad_IR") | |
| for p in (ROOT, RAD_IR_ROOT): | |
| if p not in sys.path: | |
| sys.path.insert(0, p) | |
| from radir import lorentz as L | |
| from radir.RadIR import RADIR, l2norm | |
| from radir.data.hmclip_mimic_box import HyCoClipMimicBoxDataset_JPG, collate_hmclip_box | |
| from transformer_maskgit.transformer_maskgit.ctvit import CTViT | |
| # "Normal" is handled separately; remaining are disease classes. | |
| DISEASE_DEFS = [ | |
| ("Pneumothorax", ["cls/cig/af/pneumothorax"]), | |
| ("Edema", ["cls/cig/af/pulmonary edema/hazy opacity"]), | |
| ("Pneumonia", ["cls/cig/disease/pneumonia"]), | |
| ("Effusion", ["cls/cig/af/pleural effusion"]), | |
| ("Cardiomegaly", ["cls/cig/af/enlarged cardiac silhouette"]), | |
| ("Atelectasis", ["cls/cig/af/atelectasis"]), | |
| ] | |
| NORMAL_COLS = ["cls/cig/normal"] | |
| # Used only by exclusive / priority multi-class modes | |
| ALL_CLASS_DEFS = [("Normal", NORMAL_COLS)] + DISEASE_DEFS | |
| CLASS_ORDER = [name for name, _ in ALL_CLASS_DEFS] + ["Other"] | |
| CLASS_COLORS = { | |
| # High-contrast qualitative palette (maximally separable on white) | |
| "Normal": "#00C853", # vivid green | |
| "Pneumothorax": "#AA00FF", # vivid purple | |
| "Edema": "#00B8D4", # cyan | |
| "Pneumonia": "#D50000", # strong red | |
| "Effusion": "#2962FF", # strong blue | |
| "Cardiomegaly": "#FF6D00", # strong orange | |
| "Atelectasis": "#FFD600", # gold / yellow | |
| "Other": "#212121", # near-black | |
| "Disease+": "#D50000", | |
| "Rest": "#9E9E9E", | |
| } | |
| # Image = filled circle; Text = triangle. Same hue for both; marker separates modality. | |
| MODALITY_MARKER = {"Image": "o", "Text": "^"} | |
| def lighten_hex(hex_color: str, factor: float = 0.25) -> str: | |
| """Slightly lighten a color (kept mild so points stay visible on white).""" | |
| hex_color = hex_color.lstrip("#") | |
| r, g, b = (int(hex_color[i : i + 2], 16) for i in (0, 2, 4)) | |
| r = int(r + (255 - r) * factor) | |
| g = int(g + (255 - g) * factor) | |
| b = int(b + (255 - b) * factor) | |
| return f"#{r:02x}{g:02x}{b:02x}" | |
| def style_for_class_modality(cls: str, modality: str) -> tuple[str, str]: | |
| """Same saturated color for Image/Text; marker distinguishes modality.""" | |
| color = CLASS_COLORS.get(cls, "#333333") | |
| marker = MODALITY_MARKER[modality] | |
| return color, marker | |
| def dicom_from_img_path(img_path: str) -> str: | |
| return os.path.splitext(os.path.basename(img_path.replace("\\", "/")))[0] | |
| def load_valid_dicom_ids(jsonl_path: str) -> list[str]: | |
| ids = [] | |
| with open(jsonl_path, encoding="utf-8") as f: | |
| for line in f: | |
| item = json.loads(line.strip()) | |
| ids.append(dicom_from_img_path(item["img_path"])) | |
| return ids | |
| def load_observation_table(jsonl_path: str, obs_csv_path: str): | |
| """ | |
| Align Chest ImaGenome rows to valid.jsonl order. | |
| Returns: | |
| dicom_ids: list[str] | |
| normal_mask: [N] bool — normal==1 and no interest disease | |
| disease_masks: dict[str, np.ndarray[N]] inclusive disease positives | |
| unmatched: list[(idx, dicom_id)] | |
| """ | |
| dicom_ids = load_valid_dicom_ids(jsonl_path) | |
| needed = ["dicom_id"] + NORMAL_COLS | |
| for _, cols in DISEASE_DEFS: | |
| needed.extend(cols) | |
| needed = list(dict.fromkeys(needed)) | |
| obs = pd.read_csv(obs_csv_path) | |
| missing = [c for c in needed if c not in obs.columns] | |
| if missing: | |
| raise KeyError(f"Missing columns in observations CSV: {missing}") | |
| obs = obs[needed].drop_duplicates(subset=["dicom_id"], keep="first") | |
| obs = obs.set_index("dicom_id") | |
| n = len(dicom_ids) | |
| normal_flag = np.zeros(n, dtype=bool) | |
| disease_masks = {name: np.zeros(n, dtype=bool) for name, _ in DISEASE_DEFS} | |
| unmatched = [] | |
| for i, did in enumerate(dicom_ids): | |
| if did not in obs.index: | |
| unmatched.append((i, did)) | |
| continue | |
| row = obs.loc[did] | |
| is_normal = any(float(row.get(c, 0.0)) > 0.5 for c in NORMAL_COLS) | |
| any_disease = False | |
| for name, cols in DISEASE_DEFS: | |
| hit = any(float(row.get(c, 0.0)) > 0.5 for c in cols) | |
| disease_masks[name][i] = hit | |
| any_disease = any_disease or hit | |
| # Pure normal: marked normal and none of the interest diseases | |
| normal_flag[i] = is_normal and (not any_disease) | |
| return dicom_ids, normal_flag, disease_masks, unmatched | |
| def build_multiclass_labels( | |
| normal_flag: np.ndarray, | |
| disease_masks: dict[str, np.ndarray], | |
| label_mode: str, | |
| ) -> tuple[list[str | None], dict]: | |
| """exclusive / priority single-vector labels (legacy multi-class modes).""" | |
| n = len(normal_flag) | |
| labels: list[str | None] = [] | |
| n_multi = n_none = 0 | |
| for i in range(n): | |
| hits = [] | |
| if normal_flag[i]: | |
| hits.append("Normal") | |
| for name, mask in disease_masks.items(): | |
| if mask[i]: | |
| hits.append(name) | |
| if label_mode == "exclusive": | |
| if len(hits) == 1: | |
| labels.append(hits[0]) | |
| else: | |
| labels.append(None) | |
| if len(hits) == 0: | |
| n_none += 1 | |
| else: | |
| n_multi += 1 | |
| elif label_mode == "priority": | |
| if normal_flag[i]: | |
| labels.append("Normal") | |
| else: | |
| assigned = None | |
| for name, _ in DISEASE_DEFS: | |
| if disease_masks[name][i]: | |
| assigned = name | |
| break | |
| labels.append(assigned if assigned is not None else "Other") | |
| else: | |
| raise ValueError(label_mode) | |
| kept = [l for l in labels if l is not None] | |
| stats = { | |
| "label_mode": label_mode, | |
| "n_kept": len(kept), | |
| "n_dropped_multi": n_multi, | |
| "n_dropped_none": n_none, | |
| "label_counts": dict(Counter(kept)), | |
| "n_normal_only": int(normal_flag.sum()), | |
| "disease_inclusive_counts": { | |
| k: int(v.sum()) for k, v in disease_masks.items() | |
| }, | |
| } | |
| return labels, stats | |
| def preprocess_image_like_forward(image: torch.Tensor) -> torch.Tensor: | |
| if image.ndim == 3: | |
| image = image.unsqueeze(0) | |
| if image.ndim == 4: | |
| if image.shape[1] != 1: | |
| image = image.mean(dim=1, keepdim=True) | |
| elif image.ndim == 5: | |
| if image.shape[1] != 1: | |
| image = image.mean(dim=1, keepdim=True) | |
| else: | |
| raise ValueError( | |
| f"Unexpected image ndim={image.ndim}, shape={tuple(image.shape)}" | |
| ) | |
| return image | |
| def extract_features(model, dataloader, device, modal_embedding: bool = False): | |
| all_eu_img, all_eu_txt = [], [] | |
| all_hyp_img, all_hyp_txt = [], [] | |
| model.eval() | |
| for batch in dataloader: | |
| imgs = preprocess_image_like_forward( | |
| batch["imgs"].to(device, non_blocking=True) | |
| ) | |
| enc_image_global = model.visual_transformer( | |
| imgs, | |
| return_encoded_tokens=True, | |
| modal_embedding=modal_embedding, | |
| modal_indexs=None, | |
| is_condition=False, | |
| ) | |
| if enc_image_global.dim() == 2: | |
| image_global = l2norm(enc_image_global) | |
| elif enc_image_global.dim() == 3: | |
| image_global = l2norm(enc_image_global[:, 0, :]) | |
| else: | |
| raise ValueError( | |
| f"Unexpected global visual feature shape: {enc_image_global.shape}" | |
| ) | |
| text_embeddings = model.text_transformer( | |
| input_ids=batch["caption_ids"].to(device, non_blocking=True), | |
| token_type_ids=batch["token_type_ids"].to(device, non_blocking=True), | |
| attention_mask=batch["attention_mask"].to(device, non_blocking=True), | |
| ) | |
| text_latents_global = text_embeddings[0][:, 0, :] | |
| all_eu_img.append(image_global.cpu()) | |
| all_eu_txt.append(l2norm(text_latents_global).cpu()) | |
| all_hyp_img.append(model.project_img(image_global).cpu()) | |
| all_hyp_txt.append(model.project_txt(text_latents_global).cpu()) | |
| return ( | |
| torch.cat(all_eu_img, dim=0), | |
| torch.cat(all_eu_txt, dim=0), | |
| torch.cat(all_hyp_img, dim=0), | |
| torch.cat(all_hyp_txt, dim=0), | |
| ) | |
| def prepare_image_feats( | |
| eu_img: torch.Tensor, | |
| hyp_img: torch.Tensor, | |
| feat_type: str, | |
| curv: torch.Tensor, | |
| ) -> np.ndarray: | |
| if feat_type == "hyperbolic": | |
| with torch.no_grad(): | |
| return L.log_map0(hyp_img.float(), curv.detach().float()).cpu().numpy() | |
| if feat_type == "euclidean": | |
| return eu_img.detach().cpu().numpy() | |
| raise ValueError(f"Unknown feat_type={feat_type!r}") | |
| def subsample_indices( | |
| labels: list[str], | |
| max_per_class: int | None, | |
| seed: int, | |
| ) -> np.ndarray: | |
| """Return indices into the sample list, optionally balanced by label.""" | |
| n = len(labels) | |
| idx = np.arange(n) | |
| if max_per_class is None or max_per_class <= 0: | |
| return idx | |
| rng = np.random.default_rng(seed) | |
| keep = [] | |
| for cls in sorted(set(labels)): | |
| idxs = [i for i, l in enumerate(labels) if l == cls] | |
| if len(idxs) > max_per_class: | |
| idxs = list(rng.choice(idxs, size=max_per_class, replace=False)) | |
| keep.extend(idxs) | |
| return np.sort(np.asarray(keep)) | |
| def joint_img_txt_feats( | |
| img_feats: np.ndarray, | |
| txt_feats: np.ndarray, | |
| sample_idx: np.ndarray, | |
| sample_labels: list[str], | |
| ) -> tuple[np.ndarray, list[str], list[str]]: | |
| """ | |
| Stack image then text features for the same samples. | |
| Returns feats [2M, D], class_labels [2M], modalities [2M]. | |
| """ | |
| img = img_feats[sample_idx] | |
| txt = txt_feats[sample_idx] | |
| labs = [sample_labels[i] for i in range(len(sample_idx))] | |
| feats = np.concatenate([img, txt], axis=0) | |
| class_labels = labs + labs | |
| modalities = ["Image"] * len(labs) + ["Text"] * len(labs) | |
| return feats, class_labels, modalities | |
| def run_tsne(feats: np.ndarray, perplexity: float, seed: int) -> np.ndarray: | |
| n = feats.shape[0] | |
| max_perp = max(2.0, (n - 1) / 3.0) | |
| eff_perp = min(perplexity, max_perp) | |
| tsne = TSNE( | |
| n_components=2, | |
| perplexity=eff_perp, | |
| random_state=seed, | |
| init="pca", | |
| learning_rate="auto", | |
| ) | |
| return tsne.fit_transform(feats) | |
| def scatter_class_modality( | |
| ax, | |
| emb: np.ndarray, | |
| class_labels: list[str], | |
| modalities: list[str], | |
| class_order: list[str], | |
| max_pair_lines: int = 0, | |
| ) -> None: | |
| """Color by disease/class, marker by Image vs Text.""" | |
| if max_pair_lines > 0 and len(class_labels) % 2 == 0: | |
| m = len(class_labels) // 2 | |
| n_lines = min(max_pair_lines, m) | |
| for i in range(n_lines): | |
| ax.plot( | |
| [emb[i, 0], emb[i + m, 0]], | |
| [emb[i, 1], emb[i + m, 1]], | |
| color="gray", | |
| alpha=0.12, | |
| linewidth=0.5, | |
| zorder=0, | |
| ) | |
| present = [c for c in class_order if c in set(class_labels)] | |
| present += [c for c in sorted(set(class_labels)) if c not in present] | |
| for cls in present: | |
| for modality in ("Image", "Text"): | |
| mask = np.array( | |
| [ | |
| (c == cls and m == modality) | |
| for c, m in zip(class_labels, modalities) | |
| ] | |
| ) | |
| if not mask.any(): | |
| continue | |
| color, marker = style_for_class_modality(cls, modality) | |
| ax.scatter( | |
| emb[mask, 0], | |
| emb[mask, 1], | |
| c=color, | |
| marker=marker, | |
| alpha=0.85, | |
| s=36 if modality == "Image" else 48, | |
| label=f"{cls}-{modality} (n={mask.sum()})", | |
| edgecolors="black", | |
| linewidths=0.35, | |
| zorder=2 if modality == "Image" else 3, | |
| ) | |
| def plot_disease_tsne( | |
| emb: np.ndarray, | |
| class_labels: list[str], | |
| modalities: list[str], | |
| out_path: str, | |
| title: str, | |
| dpi: int, | |
| max_pair_lines: int = 40, | |
| ) -> None: | |
| fig, ax = plt.subplots(figsize=(11, 9)) | |
| scatter_class_modality( | |
| ax, emb, class_labels, modalities, CLASS_ORDER, max_pair_lines=max_pair_lines | |
| ) | |
| ax.set_title(title) | |
| ax.set_xlabel("t-SNE 1") | |
| ax.set_ylabel("t-SNE 2") | |
| ax.legend(loc="best", frameon=True, fontsize=8) | |
| ax.grid(True, alpha=0.2) | |
| fig.tight_layout() | |
| os.makedirs(os.path.dirname(os.path.abspath(out_path)) or ".", exist_ok=True) | |
| fig.savefig(out_path, dpi=dpi, bbox_inches="tight") | |
| plt.close(fig) | |
| print(f"Saved disease t-SNE plot to {out_path}") | |
| def plot_vs_normal_panels( | |
| img_feats: np.ndarray, | |
| txt_feats: np.ndarray, | |
| normal_flag: np.ndarray, | |
| disease_masks: dict[str, np.ndarray], | |
| out_path: str, | |
| perplexity: float, | |
| seed: int, | |
| dpi: int, | |
| max_per_class: int | None, | |
| title_prefix: str, | |
| max_pair_lines: int = 30, | |
| ) -> None: | |
| """ | |
| One subplot per disease: Normal-only vs Disease+ (inclusive). | |
| Image (circle) and Text (triangle) of the same class use related colors. | |
| Joint t-SNE over image+text points in each panel. | |
| """ | |
| diseases = [name for name, _ in DISEASE_DEFS] | |
| n_panels = len(diseases) | |
| ncols = 3 | |
| nrows = int(np.ceil(n_panels / ncols)) | |
| fig, axes = plt.subplots(nrows, ncols, figsize=(5.6 * ncols, 5.0 * nrows)) | |
| axes = np.atleast_1d(axes).ravel() | |
| for ax_i, disease in enumerate(diseases): | |
| ax = axes[ax_i] | |
| pos = disease_masks[disease] | |
| keep = normal_flag | pos | |
| raw_idx = np.where(keep)[0] | |
| if raw_idx.size < 10: | |
| ax.set_title(f"{disease}: too few samples") | |
| ax.axis("off") | |
| continue | |
| sample_labels = [disease if pos[j] else "Normal" for j in raw_idx] | |
| local_keep = subsample_indices(sample_labels, max_per_class, seed + ax_i) | |
| sample_idx = raw_idx[local_keep] | |
| sample_labels = [sample_labels[i] for i in local_keep] | |
| feats, class_labels, modalities = joint_img_txt_feats( | |
| img_feats, txt_feats, sample_idx, sample_labels | |
| ) | |
| emb = run_tsne(feats, perplexity, seed + ax_i) | |
| scatter_class_modality( | |
| ax, | |
| emb, | |
| class_labels, | |
| modalities, | |
| ["Normal", disease], | |
| max_pair_lines=max_pair_lines, | |
| ) | |
| ax.set_title(f"Normal vs {disease}\n○ Image △ Text") | |
| ax.legend(loc="best", fontsize=7, frameon=True) | |
| ax.set_xticks([]) | |
| ax.set_yticks([]) | |
| ax.grid(True, alpha=0.15) | |
| for j in range(n_panels, len(axes)): | |
| axes[j].axis("off") | |
| fig.suptitle(title_prefix + " | Image○ / Text△", fontsize=13) | |
| fig.tight_layout(rect=[0, 0, 1, 0.96]) | |
| os.makedirs(os.path.dirname(os.path.abspath(out_path)) or ".", exist_ok=True) | |
| fig.savefig(out_path, dpi=dpi, bbox_inches="tight") | |
| plt.close(fig) | |
| print(f"Saved vs-normal panel t-SNE to {out_path}") | |
| def plot_highlight_panels( | |
| img_feats: np.ndarray, | |
| txt_feats: np.ndarray, | |
| disease_masks: dict[str, np.ndarray], | |
| out_path: str, | |
| perplexity: float, | |
| seed: int, | |
| dpi: int, | |
| title_prefix: str, | |
| ) -> None: | |
| """ | |
| Shared joint Image+Text t-SNE over all samples; each panel highlights one | |
| disease's image/text points against gray background. | |
| """ | |
| n = img_feats.shape[0] | |
| sample_idx = np.arange(n) | |
| sample_labels = ["Rest"] * n | |
| feats, class_labels, modalities = joint_img_txt_feats( | |
| img_feats, txt_feats, sample_idx, sample_labels | |
| ) | |
| emb = run_tsne(feats, perplexity, seed) | |
| emb_img, emb_txt = emb[:n], emb[n:] | |
| diseases = [name for name, _ in DISEASE_DEFS] | |
| ncols = 3 | |
| nrows = int(np.ceil(len(diseases) / ncols)) | |
| fig, axes = plt.subplots(nrows, ncols, figsize=(5.6 * ncols, 5.0 * nrows)) | |
| axes = np.atleast_1d(axes).ravel() | |
| for ax_i, disease in enumerate(diseases): | |
| ax = axes[ax_i] | |
| pos = disease_masks[disease] | |
| # background: non-positive image+text | |
| ax.scatter( | |
| emb_img[~pos, 0], | |
| emb_img[~pos, 1], | |
| c=CLASS_COLORS["Rest"], | |
| marker="o", | |
| alpha=0.15, | |
| s=8, | |
| edgecolors="none", | |
| label=f"Rest-Image (n={(~pos).sum()})", | |
| ) | |
| ax.scatter( | |
| emb_txt[~pos, 0], | |
| emb_txt[~pos, 1], | |
| c=lighten_hex(CLASS_COLORS["Rest"], 0.2), | |
| marker="^", | |
| alpha=0.15, | |
| s=10, | |
| edgecolors="none", | |
| label=f"Rest-Text (n={(~pos).sum()})", | |
| ) | |
| c_img, m_img = style_for_class_modality(disease, "Image") | |
| c_txt, m_txt = style_for_class_modality(disease, "Text") | |
| ax.scatter( | |
| emb_img[pos, 0], | |
| emb_img[pos, 1], | |
| c=c_img, | |
| marker=m_img, | |
| alpha=0.9, | |
| s=36, | |
| edgecolors="black", | |
| linewidths=0.35, | |
| label=f"{disease}-Image (n={pos.sum()})", | |
| ) | |
| ax.scatter( | |
| emb_txt[pos, 0], | |
| emb_txt[pos, 1], | |
| c=c_txt, | |
| marker=m_txt, | |
| alpha=0.9, | |
| s=48, | |
| edgecolors="black", | |
| linewidths=0.35, | |
| label=f"{disease}-Text (n={pos.sum()})", | |
| ) | |
| ax.set_title(f"Highlight: {disease}\n○ Image △ Text") | |
| ax.legend(loc="best", fontsize=7, frameon=True) | |
| ax.set_xticks([]) | |
| ax.set_yticks([]) | |
| ax.grid(True, alpha=0.15) | |
| for j in range(len(diseases), len(axes)): | |
| axes[j].axis("off") | |
| fig.suptitle(title_prefix + " | Image○ / Text△", fontsize=13) | |
| fig.tight_layout(rect=[0, 0, 1, 0.96]) | |
| os.makedirs(os.path.dirname(os.path.abspath(out_path)) or ".", exist_ok=True) | |
| fig.savefig(out_path, dpi=dpi, bbox_inches="tight") | |
| plt.close(fig) | |
| print(f"Saved highlight panel t-SNE to {out_path}") | |
| def build_model(args, device: str) -> RADIR: | |
| tokenizer = AutoTokenizer.from_pretrained( | |
| args.model_path, trust_remote_code=True, local_files_only=True | |
| ) | |
| text_model = AutoModel.from_pretrained( | |
| args.model_path, trust_remote_code=True, local_files_only=True | |
| ) | |
| image_encoder = CTViT( | |
| dim=768, | |
| codebook_size=8192, | |
| image_size=args.imsize, | |
| patch_size=16, | |
| temporal_patch_size=10, | |
| spatial_depth=8, | |
| temporal_depth=6, | |
| cls_depth=4, | |
| dim_head=32, | |
| heads=8, | |
| channels=1, | |
| ) | |
| model = RADIR( | |
| tokenizer=tokenizer, | |
| image_encoder=image_encoder, | |
| text_encoder=text_model, | |
| dim_text=768, | |
| dim_image=512, | |
| dim_latent=512, | |
| use_mlm=False, | |
| use_all_token_embeds=False, | |
| ) | |
| print(f"Loading checkpoint: {args.radir_ckpt}") | |
| model.load(args.radir_ckpt) | |
| model.to(device) | |
| return model | |
| def parse_args(): | |
| parser = argparse.ArgumentParser( | |
| description="Disease-colored t-SNE for RadIR/RadIR" | |
| ) | |
| parser.add_argument("--device", default="cuda", type=str) | |
| parser.add_argument( | |
| "--model_path", | |
| default=os.path.join( | |
| ROOT, "hf_models/microsoft/BiomedVLP-CXR-BERT-specialized" | |
| ), | |
| ) | |
| parser.add_argument( | |
| "--radir_ckpt", | |
| default=os.path.join(ROOT, "outputs/radir_final.pt"), | |
| ) | |
| parser.add_argument( | |
| "--dataset_root", | |
| default=os.path.join(ROOT, "dataset"), | |
| ) | |
| parser.add_argument("--caption_file", default="valid.jsonl") | |
| parser.add_argument( | |
| "--obs_csv", | |
| default=os.path.join(ROOT, "dataset/chest_imagenome-observations.csv"), | |
| ) | |
| parser.add_argument( | |
| "--cig_matched_dir", | |
| default=os.path.join(ROOT, "dataset/processed/matched"), | |
| ) | |
| parser.add_argument("--split", default="valid", type=str) | |
| parser.add_argument("--batch_size", default=16, type=int) | |
| parser.add_argument("--num_workers", default=8, type=int) | |
| parser.add_argument("--imsize", default=224, type=int) | |
| parser.add_argument("--max_words", default=128, type=int) | |
| parser.add_argument("--modal_embedding", action="store_true") | |
| parser.add_argument( | |
| "--feat_type", | |
| choices=["hyperbolic", "euclidean"], | |
| default="hyperbolic", | |
| ) | |
| parser.add_argument( | |
| "--modality", | |
| choices=["both"], | |
| default="both", | |
| help="Always joint Image+Text (kept for CLI compatibility)", | |
| ) | |
| parser.add_argument( | |
| "--label_mode", | |
| choices=["exclusive", "priority", "vs_normal", "highlight"], | |
| default="exclusive", | |
| help=( | |
| "exclusive: one figure, samples with exactly one label (recommended); " | |
| "priority: one figure, first-matching disease (keeps comorbidities); " | |
| "vs_normal: multi-panel Normal vs Disease+; " | |
| "highlight: multi-panel shared t-SNE, one disease highlighted each" | |
| ), | |
| ) | |
| parser.add_argument( | |
| "--max_samples", | |
| type=int, | |
| default=0, | |
| help="Subsample for exclusive/priority modes; 0 = all kept", | |
| ) | |
| parser.add_argument( | |
| "--max_per_class", | |
| type=int, | |
| default=0, | |
| help="For vs_normal: cap samples per class in each panel (0 = no cap)", | |
| ) | |
| parser.add_argument( | |
| "--max_pair_lines", | |
| type=int, | |
| default=30, | |
| help="Gray connectors between paired Image-Text points", | |
| ) | |
| parser.add_argument("--perplexity", type=float, default=30.0) | |
| parser.add_argument("--seed", type=int, default=42) | |
| parser.add_argument( | |
| "--out", | |
| default=os.path.join(ROOT, "outputs/tsne_disease_all.png"), | |
| ) | |
| parser.add_argument("--save_feats", default=None) | |
| parser.add_argument("--load_feats", default=None) | |
| parser.add_argument("--save_labels", default=None) | |
| parser.add_argument("--dpi", type=int, default=150) | |
| return parser.parse_args() | |
| def main(): | |
| args = parse_args() | |
| device = args.device | |
| max_samples = None if args.max_samples <= 0 else args.max_samples | |
| max_per_class = None if args.max_per_class <= 0 else args.max_per_class | |
| jsonl_path = os.path.join(args.dataset_root, args.caption_file) | |
| dicom_ids, normal_flag, disease_masks, unmatched = load_observation_table( | |
| jsonl_path, args.obs_csv | |
| ) | |
| print("===== Label matching =====") | |
| print(f" mode : {args.label_mode}") | |
| print(f" jsonl samples : {len(dicom_ids)}") | |
| print(f" unmatched : {len(unmatched)}") | |
| print(f" Normal-only : {int(normal_flag.sum())}") | |
| print(" Disease+ (inclusive):") | |
| for name, _ in DISEASE_DEFS: | |
| print(f" {name:14s} {int(disease_masks[name].sum())}") | |
| if args.save_labels: | |
| os.makedirs( | |
| os.path.dirname(os.path.abspath(args.save_labels)) or ".", | |
| exist_ok=True, | |
| ) | |
| rows = {"dicom_id": dicom_ids, "normal_only": normal_flag.astype(int)} | |
| for name, mask in disease_masks.items(): | |
| rows[f"{name}_pos"] = mask.astype(int) | |
| pd.DataFrame(rows).to_csv(args.save_labels, index=False) | |
| print(f"Saved labels to {args.save_labels}") | |
| if args.load_feats: | |
| print(f"Loading features from {args.load_feats}") | |
| data = np.load(args.load_feats, allow_pickle=True) | |
| eu_img = torch.from_numpy(data["eu_img"]) | |
| eu_txt = torch.from_numpy(data["eu_txt"]) | |
| hyp_img = torch.from_numpy(data["hyp_img"]) | |
| hyp_txt = torch.from_numpy(data["hyp_txt"]) | |
| curv = torch.tensor(float(data["curv"])) | |
| feat_type = ( | |
| str(data["feat_type"]) if "feat_type" in data.files else args.feat_type | |
| ) | |
| if isinstance(feat_type, np.ndarray): | |
| feat_type = str(feat_type.item()) | |
| else: | |
| model = build_model(args, device) | |
| curv = model.get_curv().cpu() | |
| print(f"Curvature = {curv.item():.6f}") | |
| dataset = HyCoClipMimicBoxDataset_JPG( | |
| split=args.split, | |
| tokenizer=model.tokenizer, | |
| dataset_root=args.dataset_root, | |
| cig_matched_dir=args.cig_matched_dir, | |
| caption_file=args.caption_file, | |
| imsize=args.imsize, | |
| max_words=args.max_words, | |
| transform=None, | |
| ) | |
| dataloader = DataLoader( | |
| dataset, | |
| batch_size=args.batch_size, | |
| shuffle=False, | |
| num_workers=args.num_workers, | |
| pin_memory=True, | |
| collate_fn=collate_hmclip_box, | |
| ) | |
| print(f"Extracting features from {len(dataset)} samples ...") | |
| eu_img, eu_txt, hyp_img, hyp_txt = extract_features( | |
| model, dataloader, device, modal_embedding=args.modal_embedding | |
| ) | |
| feat_type = args.feat_type | |
| if args.save_feats: | |
| os.makedirs( | |
| os.path.dirname(os.path.abspath(args.save_feats)) or ".", | |
| exist_ok=True, | |
| ) | |
| np.savez( | |
| args.save_feats, | |
| eu_img=eu_img.numpy(), | |
| eu_txt=eu_txt.numpy(), | |
| hyp_img=hyp_img.numpy(), | |
| hyp_txt=hyp_txt.numpy(), | |
| curv=curv.item(), | |
| feat_type=feat_type, | |
| ) | |
| print(f"Saved features to {args.save_feats}") | |
| n_feat = eu_img.shape[0] | |
| if n_feat != len(dicom_ids): | |
| raise RuntimeError( | |
| f"Feature count ({n_feat}) != jsonl count ({len(dicom_ids)})." | |
| ) | |
| img_feats = prepare_image_feats(eu_img, hyp_img, feat_type, curv) | |
| txt_feats = prepare_image_feats(eu_txt, hyp_txt, feat_type, curv) | |
| ckpt_name = ( | |
| os.path.basename(args.radir_ckpt) if not args.load_feats else "cached" | |
| ) | |
| title_prefix = ( | |
| f"RadIR ({args.label_mode}, {feat_type}, Image○/Text△) | {ckpt_name}" | |
| ) | |
| if args.label_mode == "vs_normal": | |
| plot_vs_normal_panels( | |
| img_feats, | |
| txt_feats, | |
| normal_flag, | |
| disease_masks, | |
| args.out, | |
| perplexity=args.perplexity, | |
| seed=args.seed, | |
| dpi=args.dpi, | |
| max_per_class=max_per_class, | |
| title_prefix=title_prefix, | |
| max_pair_lines=args.max_pair_lines, | |
| ) | |
| return | |
| if args.label_mode == "highlight": | |
| plot_highlight_panels( | |
| img_feats, | |
| txt_feats, | |
| disease_masks, | |
| args.out, | |
| perplexity=args.perplexity, | |
| seed=args.seed, | |
| dpi=args.dpi, | |
| title_prefix=title_prefix, | |
| ) | |
| return | |
| # exclusive / priority: joint Image+Text multiclass | |
| labels, stats = build_multiclass_labels( | |
| normal_flag, disease_masks, args.label_mode | |
| ) | |
| print(" multiclass kept counts:") | |
| for cls in CLASS_ORDER: | |
| if cls in stats["label_counts"]: | |
| print(f" {cls:14s} {stats['label_counts'][cls]}") | |
| keep_idx = np.array([l is not None for l in labels], dtype=bool) | |
| sample_idx = np.where(keep_idx)[0] | |
| sample_labels = [labels[i] for i in sample_idx] | |
| if max_samples is not None and max_samples < len(sample_idx): | |
| rng = np.random.default_rng(args.seed) | |
| pick = np.sort(rng.choice(len(sample_idx), size=max_samples, replace=False)) | |
| sample_idx = sample_idx[pick] | |
| sample_labels = [sample_labels[i] for i in pick] | |
| feats, class_labels, modalities = joint_img_txt_feats( | |
| img_feats, txt_feats, sample_idx, sample_labels | |
| ) | |
| print( | |
| f"t-SNE on {len(sample_idx)} samples ×2 modalities " | |
| f"(mode={args.label_mode})" | |
| ) | |
| emb = run_tsne(feats, args.perplexity, args.seed) | |
| plot_disease_tsne( | |
| emb, | |
| class_labels, | |
| modalities, | |
| args.out, | |
| title=f"{title_prefix} | n={len(sample_idx)}", | |
| dpi=args.dpi, | |
| max_pair_lines=args.max_pair_lines, | |
| ) | |
| if __name__ == "__main__": | |
| main() | |