Team Ai
Apppublic

smoothjazzuser/AI_Model_Explainability

sourceHugging Facemitupdated 1y agoView on Hugging Face
0likes
app.py170 linesDownload Raw Back to root
1import warnings2warnings.filterwarnings('ignore')3import torch, numpy as np, os4from torch import nn5from transformers import AutoModelForImageClassification, AutoConfig, AutoImageProcessor6import matplotlib.pyplot as plt7from PIL import Image8import saliency.core as saliency9import io10import gradio as gr11import PIL12 13model_choice = 314model_names = ["nvidia/mit-b0",'facebook/convnext-base-224', 'microsoft/resnet-18', 'microsoft/swin-tiny-patch4-window7-224']15model_name = model_names[model_choice]16device     = 'cuda' if torch.cuda.is_available() else 'cpu'17 18class Model(nn.Module):19    def __init__(self, MODEL_NAME=model_name):20        super().__init__()21        self.config = AutoConfig.from_pretrained(MODEL_NAME, finetuning_task="image-classification")22        self.model = AutoModelForImageClassification.from_pretrained(MODEL_NAME)23        self.class_len = self.config.num_labels24        self.id2label = self.config.id2label25        self.label2id = self.config.label2id26 27    def forward(self, x):28        if isinstance(x, np.ndarray):29            x = torch.from_numpy(x)30        if len(x.shape) == 3:31            x = x.unsqueeze(0)32        if x.shape[-1] == 3:33            x = x.permute(0, 3, 1, 2)34        x = x.to(device)35        x = self.model(x)36        return x.logits37 38def conv_layer_forward_hook(module, input, output):39    """Method from Examples_pytorch.ipynb for the gradcam library https://github.com/PAIR-code/saliency."""40    global last_conv_layer_outputs41    last_conv_layer_outputs[saliency.base.CONVOLUTION_LAYER_VALUES] = torch.movedim(output, 3, 1).detach().cpu().numpy()42def conv_layer_backward_hook(module, grad_input, grad_output):43    """Method from Examples_pytorch.ipynb for the gradcam library https://github.com/PAIR-code/saliency."""44    global last_conv_layer_outputs45    last_conv_layer_outputs[saliency.base.CONVOLUTION_OUTPUT_GRADIENTS] = torch.movedim(grad_output[0], 3, 1).detach().cpu().numpy()46 47auto_transformer, class_to_id, id_to_class, last_conv_layer, last_conv_layer_outputs = None, None, None, None, None48 49 50def swap_models(name):51    global model, auto_transformer, class_to_id, id_to_class, last_conv_layer, last_conv_layer_outputs52    auto_transformer = AutoImageProcessor.from_pretrained(name)53    model = Model(MODEL_NAME=name)54    model = model.to(device).eval()55    # register the hooks for the last convolution layer for Grad-Cam56    named_modules = dict(model.model.named_modules())57    last_conv_layer_name = None58    for name, module in named_modules.items():59        if isinstance(module, torch.nn.Conv2d):60            last_conv_layer_name = name61 62    last_conv_layer = named_modules[last_conv_layer_name]63    last_conv_layer_outputs = {}64 65    last_conv_layer.register_forward_hook(conv_layer_forward_hook)66    last_conv_layer.register_backward_hook(conv_layer_backward_hook)67    class_to_id = {v:k for k,v in model.model.config.id2label.items()}68    id_to_class = {k:v for k,v in model.model.config.id2label.items()}69 70swap_models(model_name)71 72def saliency_graph(img1, steps=25):73    img1 = auto_transformer(img1)74    img1 = np.squeeze(np.array(img1.pixel_values))75    if img1.shape[0] < img1.shape[1]:76        img1 = np.moveaxis(img1, 0, -1)77    img1 = (img1 - np.min(img1)) / (np.max(img1) - np.min(img1))78 79    class_idx_str = 'class_idx_str'80    def gradcam_call(images, call_model_args=None, expected_keys=None):81        if not isinstance(images, np.ndarray) and not isinstance(images, torch.Tensor) and not isinstance(images, PIL.Image.Image):82            # return two blank images83            im1 = np.zeros((224, 224, 3))84            im2 = np.zeros((224, 224, 3))85            return im1, im286        87        if len(images.shape) == 3:88            images = np.expand_dims(images, 0)89        images = torch.tensor(images, dtype=torch.float32)90        images = images.requires_grad_(True)91        target_class_idx = call_model_args[class_idx_str]92        y_pred = model(images)93        94        if saliency.base.INPUT_OUTPUT_GRADIENTS in expected_keys:95            out =  y_pred[:, target_class_idx]96            # move actual color channel to the 1st dimension97            #images = torch.movedim(images, 3, 1)98            grads = torch.autograd.grad(out, images, grad_outputs=torch.ones_like(out))99            grads = grads[0].detach().cpu().numpy()100            return {saliency.base.INPUT_OUTPUT_GRADIENTS: grads}101        else:102            hot = torch.zeroes_like(y_pred)103            hot[:, target_class_idx] = 1104            model.zero_grad()105            y_pred.backward(gradient=hot, retain_graph=True)106            return last_conv_layer_outputs107 108    im = img1.astype(np.float32)109    base = np.zeros(img1.shape)110 111    pred = model(torch.from_numpy(im))112    class_pred = pred.argmax(dim=1).item()113    call_model_args = {class_idx_str: class_pred}114    gradients = saliency.IntegratedGradients()115 116    s = gradients.GetSmoothedMask(im, gradcam_call, call_model_args, x_steps=steps, x_baseline=base, batch_size=25)117 118    smoothgrad_mask_grayscale = saliency.VisualizeImageGrayscale(s)119 120    with torch.no_grad():121        output = model.forward(img1)122        output = torch.nn.functional.softmax(output, dim=1)123        output = output.cpu().numpy()124    top_5 = [(id_to_class[int(i)], output[0][i]) for i in np.argsort(output)[0][-5:][::-1]]125 126 127    # Render the saliency masks.128    fig, ax = plt.subplots(1, 1, figsize=(10, 10))129    ax.barh([x[0] for x in top_5], [x[1] for x in top_5])130    ax.set_title('Top 5 Predictions')131    buf = io.BytesIO()132    fig.savefig(buf, format='jpg')133    buf.seek(0)134    fig_img = Image.open(buf)135    plt.close(fig)136    return smoothgrad_mask_grayscale, fig_img137 138# gradio Interface139def gradio_interface(img):140    smoothgrad_mask_grayscale, fig_img = saliency_graph(img, steps=20)141    return smoothgrad_mask_grayscale, fig_img142 143with gr.Blocks() as iface:144    #examples = gr.Examples(examples=["ex1.jpg", "ex2.jpg", "ex3.jpg", "ex4.jpg"], label="Examples", inputs="image", examples_per_page=4)145    gr.Markdown("This function finds the most critical pixels in an image for predicting a class by looking at the pixels models attend to. The best models will ideally make predictions by highlighting the expected object. Poorly generalizable models will often rely on environmental cues instead and forego looking at the most important pixels. Highlighting the most important pixels helps explain/build trust about whether a given model uses the correct features to make its prediction.")146    with gr.Row():147        with gr.Column():148            test_image = gr.Image(label="Input Image")149            input_btn = gr.Button("Classify image")150            model_select_dropdown = gr.Radio(model_names, label="Model to test", interactive=True)151        with gr.Column():152            output = gr.Image(label="Pixels used for classification")153            output2 = gr.Image(label="Top 5 Predictions")154 155    input_btn.click(gradio_interface, test_image, outputs=[output, output2])156    model_select_dropdown.change(swap_models, inputs=[model_select_dropdown])157    examples = gr.Examples(158        examples = [os.path.join('./', x) for x in os.listdir('./') if (x.endswith('.jpg') or x.endswith('.png') or x.endswith('.webp'))],159        inputs=gr.Image(),160        label="Examples",161        fn=gradio_interface,162        cache_examples=True,163        run_on_click=True,164        postprocess=True,165        preprocess=True,166        outputs=[output, output2])167 168 169iface.launch(share=False, inbrowser=True, server_name="0.0.0.0")170