""" Fashion Item Classifier with Dual Model Support Primary: NVIDIA optimized model (high performance) Fallback: HuggingFace HelloWorld0204/Classification-StyleWell-model Provides classification and matching capabilities for wardrobe items. """ from __future__ import annotations import os import json from typing import Any from collections import OrderedDict import numpy as np import torch from PIL import Image from transformers import AutoModelForImageClassification, AutoProcessor, pipeline DEFAULT_NVIDIA_MODEL_ID = os.getenv( "FASHION_CLASSIFIER_NVIDIA_MODEL", "nvidia/ViT-B-32-quickgelu" # Fast NVIDIA-optimized Vision Transformer ) DEFAULT_HF_MODEL_ID = os.getenv( "FASHION_CLASSIFIER_HF_MODEL", "HelloWorld0204/Classification-StyleWell-model" ) DEFAULT_CACHE_SIZE = int(os.getenv("FASHION_CLASSIFIER_CACHE_SIZE", "512")) class FashionClassifier: """ Dual-model fashion classifier with NVIDIA primary and HuggingFace fallback. Supports: - Item classification (category, type, pattern, color, fit, style) - Outfit matching between items - Confidence scoring """ def __init__( self, nvidia_model_id: str = DEFAULT_NVIDIA_MODEL_ID, hf_model_id: str = DEFAULT_HF_MODEL_ID, device: str | None = None, cache_size: int = DEFAULT_CACHE_SIZE, ) -> None: self.nvidia_model_id = nvidia_model_id self.hf_model_id = hf_model_id self.device = device or ("cuda" if torch.cuda.is_available() else "cpu") self.cache_size = cache_size self._classifier = None self._processor = None self._model = None self._backend = None self._load_attempted = False # Classification cache self._classification_cache: OrderedDict[str, dict[str, Any]] = OrderedDict() # Predefined fashion categories self._fashion_categories = { "topwear": ["shirt", "t-shirt", "blouse", "hoodie", "jacket", "blazer", "sweater", "coat"], "bottomwear": ["jeans", "trousers", "pants", "shorts", "skirt", "joggers", "leggings"], "footwear": ["sneaker", "boot", "loafer", "sandal", "heel", "shoe"], "accessories": ["bag", "belt", "watch", "cap", "scarf", "sunglasses", "jewelry"], "dress": ["dress", "gown", "jumpsuit", "romper"], } @property def backend_name(self) -> str: """Get the name of the currently loaded backend.""" self._ensure_model_loaded() return self._backend or "none" def _ensure_model_loaded(self) -> None: """Load the model on first use with fallback mechanism.""" if self._load_attempted: return self._load_attempted = True # Try NVIDIA model first if self._try_load_nvidia_model(): self._backend = "nvidia" return # Fall back to HuggingFace if self._try_load_hf_model(): self._backend = "huggingface" return self._backend = "none" print("[FashionClassifier] Failed to load both NVIDIA and HuggingFace models. Using fallback classification.") def _try_load_nvidia_model(self) -> bool: """Attempt to load NVIDIA optimized model.""" try: print(f"[FashionClassifier] Loading NVIDIA model: {self.nvidia_model_id}") # Try to load as image classification model try: self._model = AutoModelForImageClassification.from_pretrained( self.nvidia_model_id, trust_remote_code=True, ) self._processor = AutoProcessor.from_pretrained( self.nvidia_model_id, trust_remote_code=True, ) self._model.to(self.device) self._model.eval() print(f"[FashionClassifier] Successfully loaded NVIDIA model") return True except Exception: # If direct model load fails, try via pipeline self._classifier = pipeline( "image-classification", model=self.nvidia_model_id, device=0 if self.device == "cuda" else -1, ) print(f"[FashionClassifier] Successfully loaded NVIDIA model via pipeline") return True except Exception as e: print(f"[FashionClassifier] Failed to load NVIDIA model: {e}") return False def _try_load_hf_model(self) -> bool: """Attempt to load HuggingFace fallback model.""" try: print(f"[FashionClassifier] Loading HuggingFace model: {self.hf_model_id}") try: self._model = AutoModelForImageClassification.from_pretrained( self.hf_model_id, trust_remote_code=True, ) self._processor = AutoProcessor.from_pretrained( self.hf_model_id, trust_remote_code=True, ) self._model.to(self.device) self._model.eval() print(f"[FashionClassifier] Successfully loaded HuggingFace model") return True except Exception: # If direct model load fails, try via pipeline self._classifier = pipeline( "image-classification", model=self.hf_model_id, device=0 if self.device == "cuda" else -1, ) print(f"[FashionClassifier] Successfully loaded HuggingFace model via pipeline") return True except Exception as e: print(f"[FashionClassifier] Failed to load HuggingFace model: {e}") return False def classify_image(self, image: Image.Image | str) -> dict[str, Any]: """ Classify a fashion item from image. Args: image: PIL Image or URL string Returns: Dict with classification results: { "category": "topwear", "confidence": 0.95, "top_5": [{"label": "shirt", "score": 0.95}, ...], "backend": "nvidia|huggingface", "attributes": { "color": "blue", "pattern": "solid", "fit": "regular", "style": "casual" } } """ self._ensure_model_loaded() # Generate cache key if isinstance(image, str): cache_key = f"image:{image}" else: # For PIL images, use a simple hash cache_key = f"image:{id(image)}" cached = self._classification_cache.get(cache_key) if cached is not None: self._classification_cache.move_to_end(cache_key) return cached # Load image if needed if isinstance(image, str): try: from PIL import Image as PILImage image = PILImage.open(image) except Exception: return self._fallback_classification() # Classify if self._backend == "nvidia" or self._backend == "huggingface": result = self._classify_with_model(image) else: result = self._fallback_classification() # Cache result self._remember_classification(cache_key, result) return result def _classify_with_model(self, image: Image.Image) -> dict[str, Any]: """Classify image using loaded model.""" try: if self._classifier is not None: # Using pipeline predictions = self._classifier(image) return { "category": predictions[0]["label"] if predictions else "unknown", "confidence": float(predictions[0]["score"]) if predictions else 0.0, "top_5": [ {"label": p["label"], "score": float(p["score"])} for p in predictions[:5] ], "backend": self._backend, "attributes": self._infer_attributes(predictions), } elif self._model is not None and self._processor is not None: # Using direct model with torch.inference_mode(): inputs = self._processor(images=image, return_tensors="pt") inputs = {k: v.to(self.device) for k, v in inputs.items()} outputs = self._model(**inputs) logits = outputs.logits # Get top predictions probs = torch.softmax(logits, dim=-1) top_k = torch.topk(probs[0], k=5) predictions = [ { "label": self._model.config.id2label.get( idx.item(), f"class_{idx.item()}" ), "score": score.item(), } for idx, score in zip(top_k.indices, top_k.values) ] return { "category": predictions[0]["label"], "confidence": float(predictions[0]["score"]), "top_5": predictions, "backend": self._backend, "attributes": self._infer_attributes(predictions), } except Exception as e: print(f"[FashionClassifier] Classification failed: {e}") return self._fallback_classification() def classify_item(self, item: dict[str, Any]) -> dict[str, Any]: """ Classify a wardrobe item from metadata. Args: item: Wardrobe item dict with 'type', 'category', 'description', 'image_url' Returns: Classification result with category, confidence, and attributes """ # Try image classification first image_url = item.get("image_url") if image_url: try: return self.classify_image(image_url) except Exception as e: print(f"[FashionClassifier] Image classification failed: {e}") # Fall back to metadata-based classification return self._classify_from_metadata(item) def _classify_from_metadata(self, item: dict[str, Any]) -> dict[str, Any]: """Classify item based on metadata when image unavailable.""" type_str = str(item.get("type", "")).lower() category_str = str(item.get("category", "")).lower() description = item.get("description", {}) if isinstance(description, dict): desc_str = " ".join([ str(description.get("type", "")), str(description.get("category", "")), ]).lower() else: desc_str = str(description).lower() full_text = f"{type_str} {category_str} {desc_str}".lower() # Find best category match best_category = "unknown" best_match_count = 0 for category, keywords in self._fashion_categories.items(): match_count = sum(1 for kw in keywords if kw in full_text) if match_count > best_match_count: best_match_count = match_count best_category = category return { "category": best_category, "confidence": 0.7 if best_match_count > 0 else 0.3, "top_5": [ {"label": best_category, "score": 0.7 if best_match_count > 0 else 0.3} ], "backend": "metadata", "attributes": self._infer_attributes_from_metadata(item), } def match_items( self, item1: dict[str, Any] | Image.Image, item2: dict[str, Any] | Image.Image, match_threshold: float = 0.5, ) -> dict[str, Any]: """ Determine if two fashion items match well together. Args: item1: First wardrobe item or image item2: Second wardrobe item or image match_threshold: Confidence threshold for match (0-1) Returns: Dict with match result: { "match": True/False, "score": 0.85, "reason": "Colors complement well", "compatibility": { "color": 0.9, "style": 0.8, "pattern": 0.7, "fit": 0.8 } } """ # Classify both items if isinstance(item1, dict): class1 = self.classify_item(item1) else: class1 = self.classify_image(item1) if isinstance(item2, dict): class2 = self.classify_item(item2) else: class2 = self.classify_image(item2) # Calculate compatibility scores compatibility = { "category": self._category_compatibility(class1["category"], class2["category"]), "color": self._color_compatibility( class1["attributes"].get("color"), class2["attributes"].get("color"), ), "style": self._style_compatibility( class1["attributes"].get("style"), class2["attributes"].get("style"), ), "pattern": self._pattern_compatibility( class1["attributes"].get("pattern"), class2["attributes"].get("pattern"), ), "fit": self._fit_compatibility( class1["attributes"].get("fit"), class2["attributes"].get("fit"), ), } # Calculate overall match score overall_score = np.mean(list(compatibility.values())) # Determine reason reason = self._generate_match_reason(compatibility, class1, class2) return { "match": overall_score >= match_threshold, "score": float(overall_score), "reason": reason, "compatibility": {k: float(v) for k, v in compatibility.items()}, } def _infer_attributes(self, predictions: list[dict]) -> dict[str, str]: """Infer fashion attributes from predictions.""" label_str = " ".join([p.get("label", "") for p in predictions[:3]]).lower() return { "color": self._extract_attribute(label_str, ["black", "white", "blue", "red", "green", "yellow", "pink", "gray", "brown"], "neutral"), "pattern": self._extract_attribute(label_str, ["solid", "striped", "plaid", "floral", "geometric", "checkered"], "solid"), "fit": self._extract_attribute(label_str, ["slim", "regular", "loose", "oversized", "fitted"], "regular"), "style": self._extract_attribute(label_str, ["casual", "formal", "sporty", "vintage", "bohemian"], "casual"), } def _infer_attributes_from_metadata(self, item: dict[str, Any]) -> dict[str, str]: """Infer attributes from item metadata.""" metadata = json.dumps(item).lower() return { "color": self._extract_attribute(metadata, ["black", "white", "blue", "red", "green", "yellow", "pink", "gray", "brown"], "neutral"), "pattern": self._extract_attribute(metadata, ["solid", "striped", "plaid", "floral", "geometric", "checkered"], "solid"), "fit": self._extract_attribute(metadata, ["slim", "regular", "loose", "oversized", "fitted"], "regular"), "style": self._extract_attribute(metadata, ["casual", "formal", "sporty", "vintage", "bohemian"], "casual"), } def _extract_attribute(self, text: str, options: list[str], default: str) -> str: """Extract attribute from text by matching keywords.""" for option in options: if option in text: return option return default def _category_compatibility(self, cat1: str, cat2: str) -> float: """Score category compatibility (0-1).""" # Complementary categories complementary = { "topwear": ["bottomwear", "dress"], "bottomwear": ["topwear"], "footwear": ["topwear", "bottomwear", "dress"], "accessories": ["topwear", "bottomwear", "footwear", "dress"], "dress": ["footwear", "accessories"], } if cat1 == cat2: return 0.5 # Same category can work but usually not as primary match if cat1 in complementary and cat2 in complementary[cat1]: return 1.0 return 0.6 def _color_compatibility(self, color1: str | None, color2: str | None) -> float: """Score color compatibility (0-1).""" if not color1 or not color2: return 0.7 # Unknown colors get neutral score # Complementary color pairs complementary_pairs = { ("blue", "orange"), ("red", "green"), ("yellow", "purple"), } if {color1, color2} in complementary_pairs: return 1.0 # Neutral colors work with everything neutral = {"black", "white", "gray", "beige", "brown"} if color1 in neutral or color2 in neutral: return 0.85 # Same color if color1 == color2: return 0.75 return 0.65 def _style_compatibility(self, style1: str | None, style2: str | None) -> float: """Score style compatibility (0-1).""" if not style1 or not style2: return 0.7 if style1 == style2: return 0.9 # Some styles mix well mixable = { ("casual", "sporty"), ("formal", "vintage"), } if {style1, style2} in mixable: return 0.8 return 0.6 def _pattern_compatibility(self, pattern1: str | None, pattern2: str | None) -> float: """Score pattern compatibility (0-1).""" if not pattern1 or not pattern2: return 0.7 # Solid goes well with anything if pattern1 == "solid" or pattern2 == "solid": return 0.85 # Same pattern can work if pattern1 == pattern2: return 0.75 # Different patterns are riskier return 0.6 def _fit_compatibility(self, fit1: str | None, fit2: str | None) -> float: """Score fit compatibility (0-1).""" if not fit1 or not fit2: return 0.7 if fit1 == fit2: return 0.85 # Loose top with fitted bottom is good if {fit1, fit2} == {"loose", "fitted"}: return 0.9 # Different fits can still work return 0.7 def _generate_match_reason( self, compatibility: dict[str, float], class1: dict[str, Any], class2: dict[str, Any], ) -> str: """Generate human-readable match reason.""" reasons = [] if compatibility["color"] >= 0.85: reasons.append("Colors complement each other well") if compatibility["style"] >= 0.85: reasons.append("Styles match perfectly") if compatibility["pattern"] >= 0.85: reasons.append("Patterns work well together") if compatibility["fit"] >= 0.85: reasons.append("Fit proportions are balanced") if not reasons: if compatibility["category"] >= 0.85: reasons.append("Items are from complementary categories") else: reasons.append("Items are compatible") return ". ".join(reasons) def _fallback_classification(self) -> dict[str, Any]: """Return fallback classification when models fail.""" return { "category": "unknown", "confidence": 0.0, "top_5": [], "backend": "fallback", "attributes": { "color": "neutral", "pattern": "solid", "fit": "regular", "style": "casual", }, } def _remember_classification(self, cache_key: str, result: dict[str, Any]) -> None: """Store classification in cache with size limit.""" self._classification_cache[cache_key] = result self._classification_cache.move_to_end(cache_key) while len(self._classification_cache) > self.cache_size: self._classification_cache.popitem(last=False)