MCPcopy Create free account
hub / github.com/HealthX-Lab/MedCLIP-SAMv2 / main

Function main

saliency_maps/scripts/eval.py:38–72  ·  view source on GitHub ↗
(args)

Source from the content-addressed store, hash-verified

36 return results
37
38def main(args):
39 print("Loading models ...")
40 # TODO: Modify these three lines to adapt for your model
41 # For an example on a model finetuned for MSCXR, check demo https://github.com/YingWANGG/M2IB/blob/main/demo.ipynb
42 model = CLIPModel.from_pretrained("openai/clip-vit-base-patch32").to(device)
43 processor = CLIPProcessor.from_pretrained("openai/clip-vit-base-patch32")
44 tokenizer = CLIPTokenizerFast.from_pretrained("openai/clip-vit-base-patch32")
45 # TODO: Modify these two lines to adapt for your input data
46 # The annotation of Conceptual Captions is in TSV format without a header
47 # The first column is the caption and the second is the image url
48 df = pd.read_csv(args.data_path,sep='\t')
49 data = list(df.itertuples(index=False))
50 all_results = []
51 print("Evaluating ...")
52 for text, image_path in tqdm(sample(data, args.samples)):
53 # Load (from a url or a local path) and preprocess image
54 try:
55 image = Image.open(requests.get(image_path, stream=True, timeout=5).raw) if 'http' in image_path else Image.open(image_path).convert('RGB')
56 except:
57 print(f"Unable to load image at {image_path}", flush=True)
58 continue
59 image_feat = processor(images=image, return_tensors="pt")['pixel_values'].to(device) # 3*224*224
60 # Tokenize text
61 text_ids = torch.tensor([tokenizer.encode(text, add_special_tokens=True)]).to(device)
62 # Train information bottleneck on image and text
63 vmap = vision_heatmap_iba(text_ids, image_feat, model, args.vlayer, args.vbeta, args.vvar, progbar=False)
64 tmap = text_heatmap_iba(text_ids, image_feat, model, args.tlayer, args.tbeta, args.tvar, progbar=False)
65 # Evaluation
66 results = get_metrics(image_feat, vmap, text_ids, tmap, model)
67 results['image'] = image_path
68 results['text'] = text
69 all_results.append(results)
70 all_results = pd.DataFrame(all_results)
71 print(all_results.mean(axis=0), flush=True)
72 all_results.to_csv(args.output_path)
73
74if __name__ == '__main__':
75 parser = argparse.ArgumentParser('M2IB argument parser')

Callers 1

eval.pyFile · 0.70

Calls 6

vision_heatmap_ibaFunction · 0.90
text_heatmap_ibaFunction · 0.90
get_metricsFunction · 0.85
from_pretrainedMethod · 0.80
encodeMethod · 0.80
meanMethod · 0.80

Tested by

no test coverage detected