| import torch |
| import torch.nn as nn |
| from torchvision import models |
| import torchvision.models.detection as detection |
| import torchvision.transforms as transforms |
| from PIL import Image |
| import segmentation_models_pytorch as smp |
| import numpy as np |
| import cv2 |
| import gradio as gr |
| from pytorch_grad_cam import GradCAM |
| from pytorch_grad_cam.utils.model_targets import ClassifierOutputTarget |
| from pytorch_grad_cam.utils.image import show_cam_on_image |
| import os |
| import logging |
|
|
| |
| logging.basicConfig(level=logging.INFO, format='%(asctime)s - %(levelname)s - %(message)s') |
| logger = logging.getLogger(__name__) |
|
|
| |
| device = torch.device('cpu') |
|
|
| |
| seg_transform = transforms.Compose([ |
| transforms.Resize((256, 256)), |
| transforms.ToTensor(), |
| transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) |
| ]) |
|
|
| |
| preprocess = transforms.Compose([ |
| transforms.Resize((224, 224)), |
| transforms.ToTensor(), |
| transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) |
| ]) |
|
|
| |
| def process_image(input_image): |
| try: |
| logger.info("Loading U-Net model...") |
| seg_model = smp.Unet( |
| encoder_name='resnet34', |
| classes=1, |
| activation=None |
| ) |
| seg_model.load_state_dict(torch.load('best_dress_segmentation_model.pth', map_location=device)) |
| seg_model.to(device) |
| seg_model.eval() |
| logger.info("U-Net loaded") |
|
|
| logger.info("Loading ResNet18 model...") |
| class_list = ['casual_dress', 'denim_dress', 'evening_dress', 'jersey_dress', 'knitted_dress', |
| 'maxi_dress', 'occasion_dress', 'shift_dress', 'shirt_dress', 'work_dress'] |
| cnn_model = models.resnet18(weights=None) |
| num_ftrs = cnn_model.fc.in_features |
| cnn_model.fc = nn.Linear(num_ftrs, len(class_list)) |
| cnn_model.load_state_dict(torch.load('final_model.pth', map_location=device)) |
| cnn_model.to(device) |
| cnn_model.eval() |
| logger.info("ResNet18 loaded") |
|
|
| logger.info("Loading Faster R-CNN model...") |
| human_detector = detection.fasterrcnn_resnet50_fpn(weights='DEFAULT') |
| human_detector.eval().to(device) |
| logger.info("Faster R-CNN loaded") |
|
|
| logger.info("Processing image...") |
| img = Image.fromarray(input_image).convert('RGB') |
| width, height = img.size |
| input_tensor = seg_transform(img).unsqueeze(0).to(device) |
| with torch.no_grad(): |
| predictions = human_detector([input_tensor]) |
| pred = predictions[0] |
| boxes = pred['boxes'] |
| labels = pred['labels'] |
| scores = pred['scores'] |
|
|
| |
| human_boxes = [] |
| human_scores = [] |
| for box, label, score in zip(boxes, labels, scores): |
| if label == 1 and score > 0.2: |
| x1, y1, x2, y2 = box.tolist() |
| human_boxes.append([x1, y1, x2, y2]) |
| human_scores.append(score) |
|
|
| if not human_boxes: |
| logger.warning("No humans detected") |
| return "No humans detected.", "No dress segmented.", "No category predicted.", "No Grad-CAM available." |
|
|
| |
| max_score_idx = torch.argmax(torch.tensor(human_scores)).item() |
| bbox = human_boxes[max_score_idx] |
| logger.info(f"Selected bounding box: {bbox}") |
|
|
| |
| crop_img = img.crop((int(bbox[0]), int(bbox[1]), int(bbox[2]), int(bbox[3]))) |
|
|
| |
| logger.info("Running segmentation...") |
| input_tensor = seg_transform(crop_img).unsqueeze(0).to(device) |
| with torch.no_grad(): |
| seg_output = seg_model(input_tensor).sigmoid().cpu().squeeze() |
| seg_mask = (seg_output > 0.5).float() |
| if seg_mask.sum() == 0: |
| logger.warning("No dress detected") |
| return "No dress detected.", "No dress segmented.", "No category predicted.", "No Grad-CAM available." |
| seg_mask_pil = transforms.ToPILImage()(seg_mask.unsqueeze(0)) |
| seg_mask_pil = seg_mask_pil.resize(crop_img.size, Image.Resampling.BILINEAR) |
| seg_mask_np = np.array(seg_mask_pil) / 255.0 |
| crop_img_np = np.array(crop_img) |
| masked_img_np = crop_img_np * np.expand_dims(seg_mask_np, axis=2) |
| masked_img = Image.fromarray(masked_img_np.astype(np.uint8)) |
|
|
| |
| logger.info("Running classification...") |
| input_tensor = preprocess(masked_img).unsqueeze(0).to(device) |
| with torch.no_grad(): |
| output = cnn_model(input_tensor) |
| pred_class = torch.argmax(output, dim=1).item() |
| category = class_list[pred_class] |
| logger.info(f"Predicted category: {category}") |
|
|
| |
| logger.info("Generating Grad-CAM...") |
| target_layers = [cnn_model.layer4[-1]] |
| cam = GradCAM(model=cnn_model, target_layers=target_layers) |
| targets = [ClassifierOutputTarget(pred_class)] |
| grayscale_cam = cam(input_tensor=input_tensor, targets=targets)[0, :] |
| visualization = show_cam_on_image(masked_img_np / 255.0, grayscale_cam, use_rgb=True, colormap=cv2.COLORMAP_JET) |
|
|
| |
| original_with_box = np.array(img) |
| cv2.rectangle(original_with_box, (int(bbox[0]), int(bbox[1])), (int(bbox[2]), int(bbox[3])), (255, 0, 0), 2) |
|
|
| logger.info("Processing complete") |
| return Image.fromarray(original_with_box), masked_img, category, Image.fromarray(visualization) |
| except Exception as e: |
| logger.error(f"Error processing image: {str(e)}") |
| return f"Error processing image: {str(e)}", "Error.", "Error.", "Error." |
|
|
| |
| logger.info("Initializing Gradio interface...") |
| with gr.Blocks() as demo: |
| gr.Markdown("### Dress Classification Pipeline\nUpload an image to detect humans, segment dresses, classify, and visualize Grad-CAM.") |
| image_input = gr.Image(label="Upload Image") |
| original_output = gr.Image(label="Original Image with Bounding Box") |
| segmented_output = gr.Image(label="Segmented Dress") |
| category_output = gr.Textbox(label="Predicted Dress Category") |
| gradcam_output = gr.Image(label="Grad-CAM Heatmap") |
|
|
| image_input.change(process_image, image_input, [original_output, segmented_output, category_output, gradcam_output]) |
|
|
| logger.info("Launching Gradio app...") |
| demo.launch(server_name="0.0.0.0", server_port=7860) |