(args)
| 36 | return results |
| 37 | |
| 38 | def 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 | |
| 74 | if __name__ == '__main__': |
| 75 | parser = argparse.ArgumentParser('M2IB argument parser') |
no test coverage detected