Team Ai
Apppublic

andreped/vit-explainer

sourceHugging Facemitupdated 3y agoView on Hugging Face
2likes
app.py101 linesDownload Raw Back to root
1import requests2import re3 4import gradio as gr5import numpy as np6from torch import topk7from torch.nn.functional import softmax8from transformers import ViTImageProcessor, ViTForImageClassification9from transformers_interpret import ImageClassificationExplainer10 11 12def load_label_data():13    file_url = "https://gist.githubusercontent.com/yrevar/942d3a0ac09ec9e5eb3a/raw/238f720ff059c1f82f368259d1ca4ffa5dd8f9f5/imagenet1000_clsidx_to_labels.txt"14    response = requests.get(file_url)15    labels = []16    pattern = '["\'](.*?)["\']'17    for line in response.text.split('\n'):18        try:19            tmp = re.findall(pattern, line)[0]20            labels.append(tmp)21        except IndexError:22            pass23    return labels24 25 26class WebUI:27    def __init__(self):28        super().__init__()29        self.nb_classes = 1030        self.processor = ViTImageProcessor.from_pretrained('google/vit-base-patch16-224')31        self.model = ViTForImageClassification.from_pretrained('google/vit-base-patch16-224')32        self.labels = load_label_data()33    34    def run_model(self, image):35        inputs = self.processor(images=image, return_tensors="pt")36        outputs = self.model(**inputs)37        outputs = softmax(outputs.logits, dim=1)38        outputs = topk(outputs, k=self.nb_classes)39        return outputs40 41    def classify_image(self, image):42        top10 = self.run_model(image)43        return {self.labels[top10[1][0][i]]: float(top10[0][0][i]) for i in range(self.nb_classes)}44 45    def explain_pred(self, image):46        image_classification_explainer = ImageClassificationExplainer(model=self.model, feature_extractor=self.processor)47        saliency = image_classification_explainer(image)48        saliency = np.squeeze(np.moveaxis(saliency, 1, 3))49        saliency[saliency >= 0.05] = 0.0550        saliency[saliency <= -0.05] = -0.0551        saliency /= np.amax(np.abs(saliency))52        return saliency53    54    def run(self):55        examples=[56            ['https://github.com/andreped/INF1600-ai-workshop/releases/download/Examples/cat.jpg'],57            ['https://github.com/andreped/INF1600-ai-workshop/releases/download/Examples/dog.jpeg'],58        ]59        with gr.Blocks() as demo:60            with gr.Row():61                image = gr.Image(height=512)62                label = gr.Label(num_top_classes=self.nb_classes)63                saliency = gr.Image(height=512, label="saliency map", show_label=True)64 65                with gr.Column(scale=0.2, min_width=150):66                    run_btn = gr.Button("Run analysis", variant="primary", elem_id="run-button")67 68                    run_btn.click(69                        fn=lambda x: self.explain_pred(x),70                        inputs=image,71                        outputs=saliency,72                    )73 74                    run_btn.click(75                        fn=lambda x: self.classify_image(x),76                        inputs=image,77                        outputs=label,78                    )79 80                    gr.Examples(81                        examples=[82                            ['https://github.com/andreped/INF1600-ai-workshop/releases/download/Examples/cat.jpg'],83                            ['https://github.com/andreped/INF1600-ai-workshop/releases/download/Examples/dog.jpeg'],84                        ],85                        inputs=image,86                        outputs=image,87                        fn=lambda x: x,88                        cache_examples=True,89                    )90        91        demo.queue().launch(server_name="0.0.0.0", server_port=7860, share=False)92 93 94def main():95    ui = WebUI()96    ui.run()97 98 99if __name__ == "__main__":100    main()101