Harshtech1/CARZero-Replication
0
1"""2zero_shot_eval.py — CARZero zero-shot classification inference.3 4FIX 1 (KEPT) — Bypass Lightning .ckpt; instantiate from config so ViT weights5 actually load into the timm backbone instead of being silently6 dropped due to torchvision/timm key mismatch.7FIX 2 (REVERTED) — Keep ImageNet normalization. The timm ViT was pretrained with8 ImageNet stats; the CARZero checkpoint inherits that default.9FIX 3 (FIXED) — Prompt ensemble uses AVERAGING only (no subtraction). Subtracting10 the negative prompt inverted scores for some diseases.11FIX 4 (KEPT) — Improved ground truth extraction in compute_metrics.py.12"""13 14import yaml15import torch16import pandas as pd17from tqdm import tqdm18from torch.utils.data import DataLoader19from torchvision import transforms20from models.carzero_net import CARZeroLightningModel21from data.openi_dataset import OpenIDataset22from transformers import AutoTokenizer23 24 25def remap_state_dict(raw_sd: dict) -> dict:26 """Map official CARZero checkpoint keys → Lightning model attribute names."""27 new_sd = {}28 for k, v in raw_sd.items():29 if k.startswith("CARZero.fusion_module."):30 new_k = k.replace("CARZero.fusion_module.", "alignment_engine.", 1)31 elif k.startswith("CARZero.img_encoder.model."):32 new_k = k.replace("CARZero.img_encoder.model.", "vision_encoder.model.", 1)33 elif k.startswith("CARZero.text_encoder.model."):34 new_k = k.replace("CARZero.text_encoder.model.", "text_encoder.bert.", 1)35 else:36 continue # drop non-model keys37 new_sd[new_k] = v38 return new_sd39 40 41def run_evaluation():42 device = torch.device("cuda" if torch.cuda.is_available() else "cpu")43 print(f"🖥️ Device: {device}")44 45 # ── FIX 1: Fresh model from config — official ViT weights load correctly ─46 with open("configs/carzero_config.yaml") as f:47 config = yaml.safe_load(f)48 model = CARZeroLightningModel(config)49 print("✅ Fresh model skeleton instantiated (timm ViT).")50 51 # Load official pretrained weights52 checkpoint = torch.load(53 "pretrain_model/carzero_pretrained.pth",54 map_location="cpu",55 weights_only=False,56 )57 raw_sd = checkpoint.get("state_dict", checkpoint)58 mapped_sd = remap_state_dict(raw_sd)59 missing, unexpected = model.load_state_dict(mapped_sd, strict=False)60 61 if missing:62 print(f"⚠️ Missing keys ({len(missing)}): {missing[:5]}...")63 if unexpected:64 print(f"⚠️ Unexpected keys ({len(unexpected)}): {unexpected[:5]}...")65 if not missing and not unexpected:66 print("✅ State dict loaded with ZERO mismatches.")67 68 model.to(device).eval()69 70 # ── Image transform: ImageNet normalization (timm ViT default) ───────────71 # The official timm ViT uses ImageNet stats. The CARZero checkpoint inherits72 # this. Confirmed: switching to [0.5,0.5,0.5] HURTS performance.73 eval_transform = transforms.Compose([74 transforms.Resize((224, 224)),75 transforms.ToTensor(),76 transforms.Normalize(mean=[0.485, 0.456, 0.406],77 std=[0.229, 0.224, 0.225]),78 ])79 80 dataset = OpenIDataset(81 csv_file="data/carzero_cleaned_reports.csv",82 image_dir="data/images/images_normalized/",83 transform=eval_transform,84 )85 dataloader = DataLoader(dataset, batch_size=1, num_workers=2, pin_memory=True)86 tokenizer = AutoTokenizer.from_pretrained("dmis-lab/biobert-v1.1")87 88 DISEASES = ["cardiomegaly", "pleural effusion", "pneumonia", "pneumothorax"]89 90 # ── FIX 3: Prompt ensemble — AVERAGE only, no subtraction ────────────────91 # Subtracting the negative prompt ("No {disease}") inverted scores for92 # cardiomegaly and pneumonia, pushing AUROC below 0.5. Instead, we average93 # multiple positive-evidence phrasings for a more robust semantic alignment.94 TEMPLATES = [95 "There is {disease}.",96 "The patient has {disease}.",97 "Findings suggest {disease}.",98 "There is a finding of {disease} in the chest radiograph.",99 ]100 101 # ── Pre-encode all text prompts once ─────────────────────────────────────102 print(f"\n🔤 Pre-encoding {len(DISEASES) * len(TEMPLATES)} prompt embeddings...")103 text_cache = {}104 with torch.no_grad():105 for disease in DISEASES:106 for i, tmpl in enumerate(TEMPLATES):107 tokens = tokenizer(108 tmpl.format(disease=disease),109 return_tensors="pt",110 max_length=97, truncation=True, padding="max_length",111 ).to(device)112 _, tg = model.text_encoder(tokens, device=device)113 text_cache[(disease, i)] = tg # [1, D]114 115 print(f"✅ {len(text_cache)} prompt embeddings cached.")116 print(f"📊 Evaluating {len(DISEASES)} pathologies × {len(dataset)} images\n")117 118 # ── Zero-shot inference ──────────────────────────────────────────────────119 predictions = []120 with torch.no_grad():121 for images, _ in tqdm(dataloader, desc="Zero-shot inference"):122 images = images.to(device)123 img_local, img_global = model.vision_encoder(images)124 scores = {}125 126 for disease in DISEASES:127 # Average SimR score across all positive templates128 total = 0.0129 for i in range(len(TEMPLATES)):130 tg = text_cache[(disease, i)]131 sr = model.alignment_engine.decoder(132 tgt=tg.unsqueeze(1), memory=img_local,133 )134 sr = model.alignment_engine.decoder_norm(sr.squeeze(1))135 sim = model.alignment_engine.mlp_head(sr)136 total += sim.item()137 138 scores[disease] = total / len(TEMPLATES)139 140 predictions.append(scores)141 142 pd.DataFrame(predictions).to_csv("zero_shot_results.csv", index=False)143 print("\n🏆 Inference complete — results saved to zero_shot_results.csv")144 145 146if __name__ == "__main__":147 run_evaluation()148 