sidhtang/_implementation_
0
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()