Team Ai
Apppublic

sidhtang/_implementation_

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
app.py273 linesDownload Raw Back to root
1 2import cv23import numpy as np4import mediapipe as mp5import torch6import torch.nn as nn7import torchvision.transforms as transforms8from PIL import Image9import gradio as gr10from enum import Enum11import colorsys12from typing import Tuple, Dict13import torch.nn.functional as F14 15class ClothingType(Enum):16    SHIRT = "shirt"17    PANTS = "pants"18    DRESS = "dress"19    JACKET = "jacket"20 21class BodySegmentation(nn.Module):22    def __init__(self):23        super().__init__()24        # Load DeepLab v3+ for semantic segmentation25        self.model = torch.hub.load('pytorch/vision:v0.10.0', 'deeplabv3_resnet50', pretrained=True)26        self.model.eval()27        28    def forward(self, x):29        return self.model(x)['out']30 31class VirtualTryOn:32    def __init__(self):33        # Initialize MediaPipe34        self.mp_pose = mp.solutions.pose35        self.mp_holistic = mp.solutions.holistic36        self.pose = self.mp_pose.Pose(37            static_image_mode=True,38            model_complexity=2,39            min_detection_confidence=0.540        )41        self.holistic = self.mp_holistic.Holistic(42            static_image_mode=True,43            model_complexity=2,44            min_detection_confidence=0.545        )46        47        # Initialize body segmentation48        self.segmentation = BodySegmentation()49        self.device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')50        self.segmentation.to(self.device)51        52        # Image transforms53        self.transforms = transforms.Compose([54            transforms.ToTensor(),55            transforms.Normalize(mean=[0.485, 0.456, 0.406], 56                              std=[0.229, 0.224, 0.225])57        ])58    59    def get_body_segmentation(self, image: np.ndarray) -> np.ndarray:60        """61        Get precise body segmentation mask62        """63        # Prepare image for model64        pil_image = Image.fromarray(cv2.cvtColor(image, cv2.COLOR_BGR2RGB))65        input_tensor = self.transforms(pil_image).unsqueeze(0).to(self.device)66        67        # Get segmentation mask68        with torch.no_grad():69            output = self.segmentation(input_tensor)70            mask = torch.argmax(output, dim=1).squeeze().cpu().numpy()71            72        # Person class is typically index 15 in COCO dataset73        return (mask == 15).astype(np.uint8)74    75    def estimate_lighting(self, image: np.ndarray) -> Dict[str, float]:76        """77        Estimate lighting conditions from the image78        """79        # Convert to HSV80        hsv = cv2.cvtColor(image, cv2.COLOR_BGR2HSV)81        82        # Get average brightness and saturation83        brightness = np.mean(hsv[:, :, 2])84        saturation = np.mean(hsv[:, :, 1])85        86        return {87            'brightness': brightness / 255.0,88            'saturation': saturation / 255.089        }90    91    def adjust_clothing_color(self, clothing: np.ndarray, 92                            lighting_params: Dict[str, float]) -> np.ndarray:93        """94        Adjust clothing colors to match lighting conditions95        """96        # Convert to HSV for easier adjustment97        hsv = cv2.cvtColor(clothing, cv2.COLOR_BGR2HSV).astype(np.float32)98        99        # Adjust brightness and saturation100        hsv[:, :, 2] *= lighting_params['brightness']101        hsv[:, :, 1] *= lighting_params['saturation']102        103        # Ensure values are within valid range104        hsv = np.clip(hsv, 0, 255).astype(np.uint8)105        106        # Convert back to BGR107        return cv2.cvtColor(hsv, cv2.COLOR_HSV2BGR)108    109    def get_clothing_dimensions(self, landmarks, image_shape: Tuple[int, int], 110                              clothing_type: ClothingType) -> Dict:111        """112        Get clothing dimensions based on body landmarks and clothing type113        """114        height, width = image_shape[:2]115        116        if clothing_type in [ClothingType.SHIRT, ClothingType.JACKET]:117            # For upper body clothing118            left_shoulder = landmarks.landmark[self.mp_pose.PoseLandmark.LEFT_SHOULDER]119            right_shoulder = landmarks.landmark[self.mp_pose.PoseLandmark.RIGHT_SHOULDER]120            left_hip = landmarks.landmark[self.mp_pose.PoseLandmark.LEFT_HIP]121            122            shoulder_width = abs(right_shoulder.x - left_shoulder.x) * width123            torso_height = abs(left_shoulder.y - left_hip.y) * height124            125            return {126                'top_left': (127                    int(min(left_shoulder.x, right_shoulder.x) * width),128                    int(left_shoulder.y * height)129                ),130                'width': int(shoulder_width * 1.3),131                'height': int(torso_height * 1.1)132            }133            134        elif clothing_type == ClothingType.PANTS:135            # For pants136            left_hip = landmarks.landmark[self.mp_pose.PoseLandmark.LEFT_HIP]137            right_hip = landmarks.landmark[self.mp_pose.PoseLandmark.RIGHT_HIP]138            left_ankle = landmarks.landmark[self.mp_pose.PoseLandmark.LEFT_ANKLE]139            140            hip_width = abs(right_hip.x - left_hip.x) * width141            leg_height = abs(left_hip.y - left_ankle.y) * height142            143            return {144                'top_left': (145                    int(min(left_hip.x, right_hip.x) * width),146                    int(left_hip.y * height)147                ),148                'width': int(hip_width * 1.5),149                'height': int(leg_height * 1.05)150            }151            152        elif clothing_type == ClothingType.DRESS:153            # For dresses154            left_shoulder = landmarks.landmark[self.mp_pose.PoseLandmark.LEFT_SHOULDER]155            right_shoulder = landmarks.landmark[self.mp_pose.PoseLandmark.RIGHT_SHOULDER]156            left_knee = landmarks.landmark[self.mp_pose.PoseLandmark.LEFT_KNEE]157            158            shoulder_width = abs(right_shoulder.x - left_shoulder.x) * width159            dress_height = abs(left_shoulder.y - left_knee.y) * height160            161            return {162                'top_left': (163                    int(min(left_shoulder.x, right_shoulder.x) * width),164                    int(left_shoulder.y * height)165                ),166                'width': int(shoulder_width * 1.4),167                'height': int(dress_height * 1.1)168            }169    170    def try_on(self, person_image: np.ndarray, clothing_image: np.ndarray, 171               clothing_type: ClothingType) -> np.ndarray:172        """173        Enhanced try-on method with support for different clothing types174        """175        # Get body segmentation176        body_mask = self.get_body_segmentation(person_image)177        178        # Get pose landmarks179        results = self.pose.process(cv2.cvtColor(person_image, cv2.COLOR_BGR2RGB))180        if not results.pose_landmarks:181            raise ValueError("No person detected in the image")182        183        # Estimate lighting conditions184        lighting_params = self.estimate_lighting(person_image)185        186        # Adjust clothing colors187        adjusted_clothing = self.adjust_clothing_color(clothing_image, lighting_params)188        189        # Get clothing dimensions190        dimensions = self.get_clothing_dimensions(191            results.pose_landmarks, 192            person_image.shape, 193            clothing_type194        )195        196        # Resize clothing197        clothing_resized = cv2.resize(198            adjusted_clothing,199            (dimensions['width'], dimensions['height']),200            interpolation=cv2.INTER_AREA201        )202        203        # Create alpha mask for smooth blending204        if clothing_resized.shape[2] == 4:205            alpha_channel = clothing_resized[:, :, 3] / 255.0206        else:207            alpha_channel = np.ones(clothing_resized.shape[:2])208        209        alpha_3channel = np.stack([alpha_channel] * 3, axis=2)210        211        # Calculate placement coordinates212        y1 = dimensions['top_left'][1]213        y2 = y1 + dimensions['height']214        x1 = dimensions['top_left'][0]215        x2 = x1 + dimensions['width']216        217        # Ensure coordinates are within image boundaries218        y1 = max(0, y1)219        y2 = min(person_image.shape[0], y2)220        x1 = max(0, x1)221        x2 = min(person_image.shape[1], x2)222        223        # Apply body mask to improve blending224        body_mask_roi = body_mask[y1:y2, x1:x2]225        alpha_3channel = alpha_3channel * np.expand_dims(body_mask_roi, axis=2)226        227        # Blend images228        roi = person_image[y1:y2, x1:x2]229        clothing_rgb = clothing_resized[:, :, :3]230        blended = (1 - alpha_3channel) * roi + alpha_3channel * clothing_rgb[:roi.shape[0], :roi.shape[1]]231        232        result = person_image.copy()233        result[y1:y2, x1:x2] = blended234        235        return result236 237def create_gradio_interface():238    def process_images(person_img, clothing_img, clothing_type):239        try_on = VirtualTryOn()240        241        # Convert clothing type string to enum242        clothing_type_enum = ClothingType(clothing_type.lower())243        244        # Process the images245        result = try_on.try_on(person_img, clothing_img, clothing_type_enum)246        247        return result248    249    # Create the interface250    iface = gr.Interface(251        fn=process_images,252        inputs=[253            gr.Image(label="Upload Person Image"),254            gr.Image(label="Upload Clothing Image"),255            gr.Dropdown(256                choices=["Shirt", "Pants", "Dress", "Jacket"],257                label="Select Clothing Type"258            )259        ],260        outputs=gr.Image(label="Result"),261        title="Virtual Try-On System",262        description="Upload a person's image and a clothing item to see how it looks!",263        examples=[264            ["person.jpg", "shirt.png", "Shirt"],265            ["person.jpg", "pants.png", "Pants"]266        ]267    )268    269    return iface270 271if __name__ == "__main__":272    iface = create_gradio_interface()273    iface.launch()