import warnings warnings.filterwarnings("ignore", message=".*torch.meshgrid.*") from pathlib import Path from typing import List, Dict, Optional, Tuple, Union import hashlib import numpy as np import cv2 from PIL import Image import torch from SCHP import SCHP from DensePose import DensePose from mappings import * from resize_image_processor import ResizeImageProcessor from mask_utils import ( get_person_height, calculate_kernels, compute_strong_protect_area, compute_weak_protect_area, compute_mask_area, finalize_mask, keep_largest_connected_component ) class GarmentMaskProcessor: """ Processeur one-shot pour générer des masques de vêtements pour l'inpainting. Usage: processor = GarmentMaskProcessor( device="cuda:0", densepose_ckpt="path/to/densepose", schp_atr_ckpt="path/to/atr.pth", schp_lip_ckpt="path/to/lip.pth" ) results = processor.process_batch( images=["img1.jpg", "img2.jpg"], garment_types=["upper", "lower"], output_dir="./output" ) """ # Mapping des types de vêtements TYPE_TO_PART = { "upper_body": "upper", "lower_body": "lower", "full_body": "dress", "upper": "upper", "lower": "lower", "dress": "dress", "full": "dress", } # Paramètres de forme pour masques variables P_ELLIPSE = 0.50 P_BOX = 0.30 P_POLY = 0.20 def __init__( self, device: str = "cuda:0", densepose_ckpt: str = "checkpoints/DensePose", schp_atr_ckpt: str = "checkpoints/SCHP/exp-schp-201908301523-atr.pth", schp_lip_ckpt: str = "checkpoints/SCHP/exp-schp-201908261155-lip.pth", output_height: int = 1024, process_size: int = 512, use_convex_hull: bool = True, schp_batch_size: int = 12, allowed_strategies: Optional[List[str]] = None, save_images = False ): """ Initialise le processeur avec les modèles de segmentation. Args: device: Device PyTorch (ex: "cuda:0", "cpu") densepose_ckpt: Chemin vers le checkpoint DensePose schp_atr_ckpt: Chemin vers le checkpoint SCHP-ATR schp_lip_ckpt: Chemin vers le checkpoint SCHP-LIP output_height: Hauteur de sortie standardisée (défaut: 1024) process_size: Taille pour le traitement intermédiaire (défaut: 512) use_convex_hull: Appliquer convex hull au masque final schp_batch_size: Taille de batch pour SCHP (défaut: 12) allowed_strategies: Liste des stratégies autorisées parmi ["ellipse", "box", "poly"] Si None, toutes sont autorisées avec leurs probabilités par défaut """ self.output_height = output_height self.process_size = process_size self.use_convex_hull = use_convex_hull self.schp_batch_size = schp_batch_size self.save_images = save_images # Gestion des stratégies autorisées all_strategies = ["ellipse", "box", "poly"] if allowed_strategies is None: self.allowed_strategies = all_strategies else: # Validation invalid = [s for s in allowed_strategies if s not in all_strategies] if invalid: raise ValueError(f"Stratégies invalides: {invalid}. Attendu: {all_strategies}") self.allowed_strategies = list(allowed_strategies) # Recalcul des probabilités normalisées if len(self.allowed_strategies) == 0: raise ValueError("Au moins une stratégie doit être autorisée") # Probabilités par défaut default_probs = {"ellipse": self.P_ELLIPSE, "box": self.P_BOX, "poly": self.P_POLY} # Filtrer et renormaliser filtered_probs = {s: default_probs[s] for s in self.allowed_strategies} total = sum(filtered_probs.values()) self.strategy_probs = {s: p / total for s, p in filtered_probs.items()} # Initialisation des modèles if torch.cuda.is_available() and "cuda" in device: # Parse device string to extract device index if ":" in device: device_idx = int(device.split(":")[1]) else: device_idx = 0 # Default to cuda:0 if no index specified torch.cuda.set_device(device_idx) # Update device string to include index self.device = f"cuda:{device_idx}" else: self.device = device self.densepose = DensePose(densepose_ckpt, self.device) self.schp_atr = SCHP(ckpt_path=schp_atr_ckpt, device=self.device) self.schp_lip = SCHP(ckpt_path=schp_lip_ckpt, device=self.device) self.resize_processor = ResizeImageProcessor( self.densepose, default_margin_ratio=0.1, default_target_ratio=(4, 3), default_final_height=output_height, default_final_width=output_height*(3/4), ) @staticmethod def _resize_to_exact_h(img: Image.Image, target_h: int) -> Image.Image: """Redimensionne une image à une hauteur exacte en conservant le ratio.""" w, h = img.size if h == target_h: return img scale = target_h / h return img.resize((int(w * scale), target_h), Image.Resampling.LANCZOS) @staticmethod def _resize_image_fast(image: Union[str, Image.Image], target_size: int) -> Image.Image: """Redimensionne rapidement une image pour le traitement.""" if isinstance(image, str): image = Image.open(image).convert("RGB") w, h = image.size if w > h: new_w, new_h = target_size, int(h * target_size / w) else: new_h, new_w = target_size, int(w * target_size / h) return image.resize((new_w, new_h), Image.Resampling.LANCZOS) @staticmethod def _seed_from_id(id_str: str, tag: str = "VAR") -> int: """Génère un seed déterministe à partir d'un ID.""" h = hashlib.md5(f"{id_str}-{tag}".encode("utf-8")).digest() return int.from_bytes(h[:4], "little") @staticmethod def _choose_strategy(rng: np.random.Generator, strategy_probs: Dict[str, float]) -> str: """Choisit une stratégie de forme selon les probabilités données.""" strategies = list(strategy_probs.keys()) probs = list(strategy_probs.values()) return str(rng.choice(strategies, p=probs)) @staticmethod def _rho_params_from_area( mask_S: np.ndarray, H: int, W: int, person_mask: Optional[np.ndarray] = None ) -> Tuple[float, float, float]: """Calcule les paramètres de dilatation basés sur l'aire du vêtement.""" Ag = float((mask_S > 0).sum()) denom = float((person_mask > 0).sum()) if (person_mask is not None and person_mask.sum() > 0) else float(H * W) f = Ag / max(1.0, denom) if f < 0.06: rho_min, rho_max = 0.06, 0.16 rmax_frac = 0.05 elif f < 0.12: rho_min, rho_max = 0.08, 0.20 rmax_frac = 0.06 else: rho_min, rho_max = 0.10, 0.25 rmax_frac = 0.07 return rho_min, rho_max, rmax_frac @staticmethod def _dilate_to_ratio( mask_g: np.ndarray, H: int, rho_min: float, rho_max: float, rmax_frac: float, rng: np.random.Generator ) -> np.ndarray: """Dilatation elliptique pour atteindre un ratio d'anneau cible.""" Ag = float((mask_g > 0).sum()) if Ag < 1: return mask_g.copy() rho = rho_min + (rho_max - rho_min) * 0.6 r_max = int(round(rmax_frac * H)) r_min = 1 r = max(r_min, int(round(0.5 * r_max))) best = None for _ in range(8): k = max(1, 2 * r + 1) ker = cv2.getStructuringElement(cv2.MORPH_ELLIPSE, (k, k)) dil = cv2.dilate(mask_g, ker, iterations=1) rr = (float((dil > 0).sum()) - Ag) / Ag if (best is None) or (abs(rr - rho) < abs(best[0] - rho)): best = (rr, r, dil) if rr < rho * 0.98 and r < r_max: r += 1 elif rr > rho * 1.02 and r > r_min: r -= 1 else: break return best[2].copy() @staticmethod def _bbox_to_ratio( mask_g: np.ndarray, H: int, W: int, rho_min: float, rho_max: float, rng: np.random.Generator, jitter_bounds: Tuple[float, float] = (0.04, 0.30) ) -> np.ndarray: """Boîte englobante jitterée pour approximer un ratio cible.""" Ag = float((mask_g > 0).sum()) if Ag < 1: return mask_g.copy() ys, xs = np.where(mask_g > 0) if len(xs) == 0: return mask_g.copy() rho = rho_min + (rho_max - rho_min) * float(rng.beta(1.5, 6.0)) x1, x2 = int(xs.min()), int(xs.max()) y1, y2 = int(ys.min()), int(ys.max()) w = x2 - x1 + 1 h = y2 - y1 + 1 jmin, jmax = jitter_bounds best = None for _ in range(12): j = float(rng.uniform(jmin, jmax)) dx = int(round(j * w)) dy = int(round(j * h)) xa = max(0, x1 - dx); xb = min(W - 1, x2 + dx) ya = max(0, y1 - dy); yb = min(H - 1, y2 + dy) box_area = float((xb - xa + 1) * (yb - ya + 1)) rr = (box_area - Ag) / Ag out = np.zeros((H, W), np.uint8) out[ya:yb + 1, xa:xb + 1] = 255 if (best is None) or (abs(rr - rho) < abs(best[0] - rho)): best = (rr, j, out) if rr < rho * 0.98: jmin = max(jmin, j) elif rr > rho * 1.02: jmax = min(jmax, j) else: break return best[2].copy() @staticmethod def _polygonize( mask: np.ndarray, rng: np.random.Generator, eps_frac_range: Tuple[float, float] = (0.03, 0.07) ) -> np.ndarray: """Simplifie le contour en polygone.""" cnts, _ = cv2.findContours((mask > 0).astype(np.uint8), cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE) out = np.zeros_like(mask, dtype=np.uint8) for c in cnts: per = cv2.arcLength(c, True) eps = per * float(rng.uniform(*eps_frac_range)) approx = cv2.approxPolyDP(c, eps, True) cv2.fillPoly(out, [approx], 255) return out if out.any() else mask.copy() # def _trim_edges(mask: np.ndarray, trim_frac_bottom=0.05, trim_frac_sides=0.03) -> np.ndarray: # H, W = mask.shape # # Pieds : enlever une bande en bas # bh = int(H * trim_frac_bottom) # mask[H - bh:H, :] = 0 # # Mains : enlever une petite bande gauche/droite # sw = int(W * trim_frac_sides) # mask[:, :sw] = 0 # mask[:, W - sw:] = 0 # return mask def _build_variable_mask( self, mask_S: np.ndarray, H: int, W: int, strong_out: Optional[np.ndarray], seed: int, person_mask: Optional[np.ndarray] = None ) -> Tuple[np.ndarray, str, Tuple[float, float, float]]: """Construit le masque variable d'inpainting.""" rng = np.random.default_rng(seed) rho_min, rho_max, rmax_frac = self._rho_params_from_area(mask_S, H, W, person_mask) strategy = self._choose_strategy(rng, self.strategy_probs) if strategy == "ellipse": m = self._dilate_to_ratio(mask_S, H, rho_min, rho_max, rmax_frac, rng) elif strategy == "box": m = self._bbox_to_ratio(mask_S, H, W, rho_min, rho_max, rng) else: # "poly" tmp = self._dilate_to_ratio(mask_S, H, rho_min, rho_max, rmax_frac, rng) m = self._polygonize(tmp, rng) # m = self._trim_edges(m, trim_frac_bottom=0.05, trim_frac_sides=0.03) if strong_out is not None: m = cv2.bitwise_and(m, cv2.bitwise_not(strong_out)) m = keep_largest_connected_component(m) m = cv2.bitwise_or(m, mask_S) return m, strategy, (rho_min, rho_max, rmax_frac) def process_batch( self, images: List[Union[str, Image.Image]], garment_types: Union[str, List[str]], image_ids: Optional[List[str]] = None, output_dir: Optional[str] = None, save_mask_s: bool = False, save_strong_protect: bool = False ) -> List[Dict]: """ Traite un batch d'images. Args: images: Liste de chemins d'images ou d'objets PIL Image garment_types: Type(s) de vêtement à générer pour CHAQUE image - str unique: un seul type appliqué à toutes les images - List: tous ces types seront générés pour chaque image image_ids: IDs optionnels pour chaque image (pour seed déterministe) output_dir: Dossier de sortie optionnel save_mask_s: Sauvegarder le masque serré save_strong_protect: Sauvegarder les zones protégées Returns: Liste de dicts, un par image: [{ "image_id": "id1", "image_standardized": PIL.Image (RGB), "strong_protect": PIL.Image (L), "masks": { "upper": { "person_mask": PIL.Image (L), "mask_S": PIL.Image (L), "strategy": "ellipse", "rho_params": (0.1, 0.25, 0.07) }, ... } }, ...] """ # Normaliser garment_types en liste if isinstance(garment_types, str): garment_types_list = [garment_types] else: garment_types_list = list(garment_types) # Générer IDs si non fournis if image_ids is None: image_ids = [None] * len(images) if len(image_ids) != len(images): raise ValueError(f"image_ids length ({len(image_ids)}) != images length ({len(images)})") # Préparer output_dir out_path = Path(output_dir) if output_dir else None if out_path: out_path.mkdir(parents=True, exist_ok=True) results = [] # Pour chaque image for img, img_id in zip(images, image_ids): # 1) Charger et préparer l'image UNE SEULE FOIS if isinstance(img, str): person_img = Image.open(img).convert("RGB") else: person_img = img.convert("RGB") img_resized = self.resize_processor(person_img) # img_resized = self._resize_image_fast(person_img, self.process_size) # 2) Inférences des modèles UNE SEULE FOIS dp_np = np.array(self.densepose(img_resized)) atr_result = self.schp_atr([img_resized]) lip_result = self.schp_lip([img_resized]) atr_img = atr_result[0] if isinstance(atr_result, list) else atr_result lip_img = lip_result[0] if isinstance(lip_result, list) else lip_result atr_np = np.array(atr_img) lip_np = np.array(lip_img) # 3) Standardiser l'image # person_img_std = self._resize_to_exact_h(person_img, self.output_height) person_img_std = img_resized H, W = person_img_std.size[1], person_img_std.size[0] # 4) Calculer zones protégées (communes) person_height = get_person_height(dp_np) dilate_kernel, _ = calculate_kernels(person_height) strong = compute_strong_protect_area(dp_np, lip_np, atr_np, dilate_kernel) strong_out = (np.array(strong).astype(np.uint8) > 0).astype(np.uint8) * 255 if strong_out.shape[:2] != (H, W): strong_out = cv2.resize(strong_out, (W, H), interpolation=cv2.INTER_NEAREST) # self.strong_protect_dilation=3 # if self.strong_protect_dilation > 0: # extra_kernel = cv2.getStructuringElement( # cv2.MORPH_ELLIPSE, # (self.strong_protect_dilation, self.strong_protect_dilation) # ) # strong_out = cv2.dilate(strong_out, extra_kernel, iterations=1) # 5) Générer l'ID final_img_id = img_id or hashlib.md5(str(img).encode()).hexdigest()[:8] # 6) Préparer le résultat pour cette image image_result = { "image_id": final_img_id, "image_standardized": person_img_std, # "strong_protect": Image.fromarray(strong_out, mode="L"), "masks": {} } # 7) Sauvegarder l'image commune si output_dir if out_path and self.save_images: img_dir = out_path / final_img_id img_dir.mkdir(parents=True, exist_ok=True) person_img_std.save(img_dir / "person.png") if save_strong_protect: Image.fromarray(strong_out, mode="L").save(img_dir / "strong_protect.png") # 8) Générer les masques pour CHAQUE type de vêtement for gtype in garment_types_list: try: part = self.TYPE_TO_PART.get(gtype, "dress") kernal_size = calculate_kernels(person_height)[1] # Calcul du masque pour ce type weak = compute_weak_protect_area(lip_np, atr_np, strong, part) mask_area, _ = compute_mask_area(dp_np, lip_np, atr_np, weak, strong, dilate_kernel, part) final_mask_img = finalize_mask(mask_area, kernal_size, strong, strong, dilate_kernel) # Nettoyage mask_np = np.array(final_mask_img.convert("L")) mask_np = (mask_np > 127).astype(np.uint8) * 255 mask_np = keep_largest_connected_component(mask_np) mask_np = cv2.morphologyEx(mask_np, cv2.MORPH_CLOSE, dilate_kernel, iterations=1) if self.use_convex_hull: cnts, _ = cv2.findContours((mask_np > 0).astype(np.uint8), cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE) hull = np.zeros_like(mask_np, dtype=np.uint8) for c in cnts: h = cv2.convexHull(c) cv2.drawContours(hull, [h], -1, 255, thickness=-1) mask_np = hull # Redimensionner au format standard if mask_np.shape[:2] != (H, W): mask_np = cv2.resize(mask_np, (W, H), interpolation=cv2.INTER_NEAREST) mask_np = (mask_np > 127).astype(np.uint8) * 255 # Masque serré mask_S = cv2.bitwise_and(mask_np, cv2.bitwise_not(strong_out)) mask_S = keep_largest_connected_component(mask_S) # Masque variable avec seed unique seed_var = self._seed_from_id(f"{final_img_id}_{gtype}", "VAR") var_mask, strategy, rho_params = self._build_variable_mask( mask_S, H, W, strong_out=strong_out, seed=seed_var ) # Sauvegarder si demandé if out_path and self.save_images: type_dir = out_path / final_img_id / gtype type_dir.mkdir(parents=True, exist_ok=True) Image.fromarray(var_mask, mode="L").save(type_dir / "person_mask.png") if save_mask_s: Image.fromarray(mask_S, mode="L").save(type_dir / "mask_S.png") # Stocker dans le résultat image_result["masks"][gtype] = { "person_mask": Image.fromarray(var_mask, mode="L"), # "mask_S": Image.fromarray(mask_S, mode="L"), # "strategy": strategy, # "rho_params": rho_params } except Exception as e: print(f"Erreur sur l'image {final_img_id} avec type {gtype}: {e}") image_result["masks"][gtype] = {"error": str(e)} results.append(image_result) return results def __call__(self, *args, **kwargs): """Permet d'utiliser l'instance comme une fonction.""" return self.process_batch(*args, **kwargs)