Spaces:
Runtime error
Runtime error
| import cv2 | |
| import torch | |
| import gradio as gr | |
| import numpy as np | |
| from PIL import Image | |
| import torch.nn.functional as F | |
| from model import DINOv2DPT, compute_CD_heatmap, cd_to_tensor | |
| DEVICE = "cuda" if torch.cuda.is_available() else "cpu" | |
| IMG_SIZE = 518 | |
| CHECKPOINT = "best.pt" | |
| # ========================= | |
| # LOAD MODEL | |
| # ========================= | |
| model = DINOv2DPT( | |
| img_size=IMG_SIZE, | |
| features=256, | |
| hook_indices=(2, 5, 8, 11), | |
| unfreeze_blocks=3 | |
| ) | |
| state_dict = torch.load(CHECKPOINT, map_location=DEVICE) | |
| model.load_state_dict(state_dict) | |
| model.to(DEVICE) | |
| model.eval() | |
| # ========================= | |
| # INFERENCE | |
| # ========================= | |
| def predict(image): | |
| image_np = np.array(image).astype(np.uint8) | |
| original_h, original_w = image_np.shape[:2] | |
| resized = cv2.resize(image_np, (IMG_SIZE, IMG_SIZE)) | |
| input_tensor = ( | |
| torch.tensor(resized / 255.0, dtype=torch.float32) | |
| .permute(2, 0, 1) | |
| .unsqueeze(0) | |
| .to(DEVICE) | |
| ) | |
| # CD MAP | |
| image_bgr = cv2.cvtColor(resized, cv2.COLOR_RGB2BGR) | |
| cd_map = compute_CD_heatmap(image_bgr) | |
| cd_tensor = cd_to_tensor(cd_map, DEVICE) | |
| # MODEL | |
| with torch.no_grad(): | |
| logits = model(input_tensor, cd_map=cd_tensor) | |
| probs = torch.sigmoid(logits)[0, 0] | |
| mask = probs.cpu().numpy() | |
| mask = cv2.resize(mask, (original_w, original_h)) | |
| binary_mask = (mask > 0.5).astype(np.uint8) | |
| # VISUALIZATION | |
| mask_vis = (binary_mask * 255).astype(np.uint8) | |
| overlay = image_np.copy() | |
| overlay[binary_mask == 1] = [255, 0, 0] | |
| overlay = cv2.addWeighted(image_np, 0.7, overlay, 0.3, 0) | |
| return ( | |
| # image_np, | |
| mask_vis, | |
| overlay | |
| ) | |
| # ========================= | |
| # EXAMPLES | |
| # ========================= | |
| examples = [ | |
| "examples/image.png", | |
| "examples/0008.jpg", | |
| "examples/apricot.jpg", | |
| ] | |
| # ========================= | |
| # UI | |
| # ========================= | |
| with gr.Blocks(theme=gr.themes.Soft()) as demo: | |
| gr.Markdown( | |
| """ | |
| # Semantic Transformer based Adversarial Patch Segmentation | |
| Upload an image to detect adversarial patches using DINOv2 + DPT. | |
| """ | |
| ) | |
| with gr.Row(): | |
| # ===================================== | |
| # LEFT PANEL | |
| # ===================================== | |
| with gr.Column(scale=1): | |
| input_image = gr.Image( | |
| type="pil", | |
| label="Input Image", | |
| height=400 | |
| ) | |
| run_btn = gr.Button( | |
| "Run Detection", | |
| variant="primary" | |
| ) | |
| gr.Examples( | |
| examples=examples, | |
| inputs=input_image, | |
| ) | |
| # ===================================== | |
| # RIGHT PANEL | |
| # ===================================== | |
| with gr.Column(scale=2): | |
| with gr.Tab("Overlay"): | |
| overlay_output = gr.Image( | |
| label="Patch Detection Overlay", | |
| height=700 | |
| ) | |
| with gr.Tab("Binary Mask"): | |
| mask_output = gr.Image( | |
| label="Predicted Mask", | |
| height=700 | |
| ) | |
| run_btn.click( | |
| fn=predict, | |
| inputs=input_image, | |
| outputs=[ | |
| mask_output, | |
| overlay_output | |
| ] | |
| ) | |
| demo.queue().launch() |