CoolFace
Apppublic

SV12/ERA_Session13

sourceHugging Facemitupdated 2y agoView on Hugging Face
0likes
app.py289 linesDownload Raw Back to root
1import gradio as gr2import random3import numpy as np4from PIL import Image5import torch6import torchvision7 8from pytorch_grad_cam import GradCAM9from pytorch_grad_cam.utils.image import show_cam_on_image10 11from models.resnet_lightning import ResNet12from utils.data import CIFARDataModule13from utils.transforms import test_transform14from utils.common import get_misclassified_data15 16inv_normalize = torchvision.transforms.Normalize(17    mean=[-0.50 / 0.23, -0.50 / 0.23, -0.50 / 0.23], std=[1 / 0.23, 1 / 0.23, 1 / 0.23]18)19 20datamodule = CIFARDataModule()21datamodule.setup()22classes = datamodule.train_dataset.classes23 24model = ResNet.load_from_checkpoint("model.ckpt")25model = model.to("cpu")26 27prediction_image = None28 29 30def upload_file(files):31    file_paths = [file.name for file in files]32    return file_paths33 34 35def read_image(path):36    img = Image.open(path)37    img.load()38    data = np.asarray(img, dtype="uint8")39    return data40 41 42def sample_images():43    images = []44    length = len(datamodule.test_dataset)45    classes = datamodule.train_dataset.classes46    for i in range(10):47        idx = random.randint(0, length - 1)48        image, label = datamodule.test_dataset[idx]49        image = inv_normalize(image).permute(1, 2, 0).numpy()50        images.append((image, classes[label]))51    return images52 53 54def get_misclassified_images(misclassified_count):55    misclassified_images = []56    misclassified_data = get_misclassified_data(57        model=model,58        device="cpu",59        test_loader=datamodule.test_dataloader(),60        count=misclassified_count,61    )62    for i in range(misclassified_count):63        img = misclassified_data[i][0].squeeze().to("cpu")64        img = inv_normalize(img)65        img = np.transpose(img.numpy(), (1, 2, 0))66        label = f"Label: {classes[misclassified_data[i][1].item()]} | Prediction: {classes[misclassified_data[i][2].item()]}"67        misclassified_images.append((img, label))68    return misclassified_images69 70 71def get_gradcam_images(gradcam_layer, gradcam_count, gradcam_opacity):72    gradcam_images = []73    if gradcam_layer == "Layer1":74        target_layers = [model.layer1[-1]]75    elif gradcam_layer == "Layer2":76        target_layers = [model.layer2[-1]]77    else:78        target_layers = [model.layer3[-1]]79 80    cam = GradCAM(model=model, target_layers=target_layers, use_cuda=False)81    data = get_misclassified_data(82        model=model,83        device="cpu",84        test_loader=datamodule.test_dataloader(),85        count=gradcam_count,86    )87    for i in range(gradcam_count):88        input_tensor = data[i][0]89 90        # Get the activations of the layer for the images91        grayscale_cam = cam(input_tensor=input_tensor, targets=None)92        grayscale_cam = grayscale_cam[0, :]93 94        # Get back the original image95        img = input_tensor.squeeze(0).to("cpu")96        if inv_normalize is not None:97            img = inv_normalize(img)98        rgb_img = np.transpose(img, (1, 2, 0))99        rgb_img = rgb_img.numpy()100 101        # Mix the activations on the original image102        visualization = show_cam_on_image(103            rgb_img, grayscale_cam, use_rgb=True, image_weight=gradcam_opacity104        )105        label = f"Label: {classes[data[i][1].item()]} | Prediction: {classes[data[i][2].item()]}"106        gradcam_images.append((visualization, label))107    return gradcam_images108 109 110def show_hide_misclassified(status):111    if not status:112        return {misclassified_count: gr.update(visible=False)}113    return {misclassified_count: gr.update(visible=True)}114 115 116def show_hide_gradcam(status):117    if not status:118        return [gr.update(visible=False) for i in range(3)]119    return [gr.update(visible=True) for i in range(3)]120 121 122def set_prediction_image(evt: gr.SelectData, gallery):123    global prediction_image124    if isinstance(gallery[evt.index], dict):125        prediction_image = gallery[evt.index]["name"]126    else:127        prediction_image = gallery[evt.index][0]["name"]128 129 130def predict(131    is_misclassified,132    misclassified_count,133    is_gradcam,134    gradcam_count,135    gradcam_layer,136    gradcam_opacity,137    num_classes,138):139    misclassified_images = None140    if is_misclassified:141        misclassified_images = get_misclassified_images(int(misclassified_count))142 143    gradcam_images = None144    if is_gradcam:145        gradcam_images = get_gradcam_images(146            gradcam_layer, int(gradcam_count), gradcam_opacity147        )148 149    img = read_image(prediction_image)150    image_transformed = test_transform(image=img)["image"]151    output = model(image_transformed.unsqueeze(0))152    preds = torch.softmax(output, dim=1).squeeze().detach().numpy()153    indices = (154        output.argsort(descending=True).squeeze().detach().numpy()[: int(num_classes)]155    )156    predictions = {classes[i]: round(float(preds[i]), 2) for i in indices}157 158    return {159        miscalssfied_output: gr.update(value=misclassified_images),160        gradcam_output: gr.update(value=gradcam_images),161        prediction_label: gr.update(value=predictions),162    }163 164 165with gr.Blocks() as app:166    gr.Markdown("## ERA Session13 - CIFAR10 Classification with ResNet")167    with gr.Row():168        with gr.Column():169            with gr.Group():170                is_misclassified = gr.Checkbox(171                    label="Misclassified Images", info="Display misclassified images?"172                )173                misclassified_count = gr.Dropdown(174                    choices=["10", "20"],175                    label="Select Number of Images",176                    info="Number of Misclassified images",177                    visible=False,178                    interactive=True,179                )180                is_misclassified.input(181                    show_hide_misclassified,182                    inputs=[is_misclassified],183                    outputs=[misclassified_count],184                )185            with gr.Group():186                is_gradcam = gr.Checkbox(187                    label="GradCAM Images",188                    info="Display GradCAM images?",189                )190                gradcam_count = gr.Dropdown(191                    choices=["10", "20"],192                    label="Select Number of Images",193                    info="Number of GradCAM images",194                    interactive=True,195                    visible=False,196                )197                gradcam_layer = gr.Dropdown(198                    choices=["Layer1", "Layer2", "Layer3"],199                    label="Select the layer",200                    info="Please select the layer for which the GradCAM is required",201                    interactive=True,202                    visible=False,203                )204                gradcam_opacity = gr.Slider(205                    minimum=0,206                    maximum=1,207                    value=0.6,208                    label="Opacity",209                    info="Opacity of GradCAM output",210                    interactive=True,211                    visible=False,212                )213 214                is_gradcam.input(215                    show_hide_gradcam,216                    inputs=[is_gradcam],217                    outputs=[gradcam_count, gradcam_layer, gradcam_opacity],218                )219            with gr.Group():220                # file_output = gr.File(file_types=["image"])221                with gr.Group():222                    upload_gallery = gr.Gallery(223                        value=None,224                        label="Uploaded images",225                        show_label=False,226                        elem_id="gallery_upload",227                        columns=5,228                        rows=2,229                        height="auto",230                        object_fit="contain",231                    )232                    upload_button = gr.UploadButton(233                        "Click to Upload images",234                        file_types=["image"],235                        file_count="multiple",236                    )237                    upload_button.upload(upload_file, upload_button, upload_gallery)238 239                with gr.Group():240                    sample_gallery = gr.Gallery(241                        value=sample_images,242                        label="Sample images",243                        show_label=True,244                        elem_id="gallery_sample",245                        columns=5,246                        rows=2,247                        height="auto",248                        object_fit="contain",249                    )250 251                upload_gallery.select(set_prediction_image, inputs=[upload_gallery])252                sample_gallery.select(set_prediction_image, inputs=[sample_gallery])253 254            with gr.Group():255                num_classes = gr.Dropdown(256                    choices=[str(i + 1) for i in range(10)],257                    label="Select Number of Top Classes",258                    info="Number of Top target classes to be shown",259                )260            run_btn = gr.Button()261        with gr.Column():262            with gr.Group():263                miscalssfied_output = gr.Gallery(264                    value=None, label="Misclassified Images", show_label=True265                )266            with gr.Group():267                gradcam_output = gr.Gallery(268                    value=None, label="GradCAM Images", show_label=True269                )270            with gr.Group():271                prediction_label = gr.Label(value=None, label="Predictions")272 273        run_btn.click(274            predict,275            inputs=[276                is_misclassified,277                misclassified_count,278                is_gradcam,279                gradcam_count,280                gradcam_layer,281                gradcam_opacity,282                num_classes,283            ],284            outputs=[miscalssfied_output, gradcam_output, prediction_label],285        )286 287 288app.launch(server_name="0.0.0.0", server_port=8000)289