Team Ai
Modelpublic

Harshtech1/CARZero-Replication

sourceHugging Faceupdated 4mo agoView on Hugging Face
0likes
zero_shot_eval.py148 linesDownload Raw Back to root
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