Spaces:
Sleeping
Sleeping
| from __future__ import annotations | |
| import hashlib | |
| import io | |
| import json | |
| import os | |
| from collections import OrderedDict | |
| from typing import Any | |
| from urllib.parse import urlparse | |
| from urllib.request import Request, urlopen | |
| import numpy as np | |
| import torch | |
| from PIL import Image | |
| from transformers import AutoModel, AutoProcessor | |
| from .schemas import EncodedWardrobeItem, RecommendationContext, SlotName | |
| try: | |
| import open_clip | |
| except ImportError: | |
| open_clip = None | |
| DEFAULT_ENCODER_MODEL_ID = os.getenv("FASHION_ENCODER_MODEL_ID", "patrickjohncyh/fashion-clip") | |
| DEFAULT_EMBEDDING_DIM = int(os.getenv("FASHION_EMBEDDING_DIM", "512")) | |
| DEFAULT_IMAGE_TIMEOUT_SECONDS = int(os.getenv("FASHION_IMAGE_TIMEOUT_SECONDS", "8")) | |
| DEFAULT_CACHE_SIZE = int(os.getenv("FASHION_EMBEDDING_CACHE_SIZE", "2048")) | |
| _SLOT_KEYWORDS = { | |
| "top": [ | |
| "topwear", "shirt", "t-shirt", "tee", "blouse", | |
| "hoodie", "jacket", "blazer", "sweater", "polo", "coat", | |
| ], | |
| "bottom": [ | |
| "bottomwear", "jeans", "trouser", "trousers", "pant", "pants", | |
| "shorts", "skirt", "jogger", "palazzo", "leggings", "chinos", | |
| ], | |
| "shoes": [ | |
| "footwear", "shoe", "shoes", "sneaker", "boot", "loafer", | |
| "sandal", "heel", | |
| ], | |
| "accessory": [ | |
| "accessory", "accessories", "bag", "belt", "watch", "cap", | |
| "hat", "scarf", "sunglasses", "jewelry", | |
| ], | |
| } | |
| _STANDALONE_OUTFIT_KEYWORDS = [ | |
| "others", "kurta", "dress", "jumpsuit", "romper", "gown", | |
| "saree", "lehenga", "co-ord", "coord", "one-piece", "one piece", | |
| ] | |
| def infer_slot_name(item: dict[str, Any]) -> SlotName: | |
| description = item.get("description") if isinstance(item.get("description"), dict) else {} | |
| raw = " ".join( | |
| [ | |
| str(item.get("type") or ""), | |
| str(item.get("category") or ""), | |
| str(description.get("type") or ""), | |
| str(description.get("category") or ""), | |
| ] | |
| ).lower() | |
| if any(keyword in raw for keyword in _STANDALONE_OUTFIT_KEYWORDS): | |
| return "unknown" | |
| for slot, keywords in _SLOT_KEYWORDS.items(): | |
| if any(keyword in raw for keyword in keywords): | |
| return slot # type: ignore[return-value] | |
| return "unknown" | |
| class FashionItemEncoder: | |
| """ | |
| Multimodal garment encoder. | |
| Output shape: | |
| encode_item(...).vector -> [D] | |
| encode_context(...) -> [D] | |
| """ | |
| def __init__( | |
| self, | |
| model_id: str = DEFAULT_ENCODER_MODEL_ID, | |
| embedding_dim: int = DEFAULT_EMBEDDING_DIM, | |
| device: str | None = None, | |
| image_timeout_seconds: int = DEFAULT_IMAGE_TIMEOUT_SECONDS, | |
| cache_size: int = DEFAULT_CACHE_SIZE, | |
| ) -> None: | |
| self.model_id = model_id | |
| self.embedding_dim = embedding_dim | |
| self.device = device or ("cuda" if torch.cuda.is_available() else "cpu") | |
| self.image_timeout_seconds = image_timeout_seconds | |
| self.cache_size = cache_size | |
| self._backend = "fallback-hash" | |
| self._model = None | |
| self._processor = None | |
| self._preprocess = None | |
| self._tokenizer = None | |
| self._load_attempted = False | |
| self._embedding_cache: OrderedDict[str, np.ndarray] = OrderedDict() | |
| def backend_name(self) -> str: | |
| self._ensure_model_loaded() | |
| return self._backend | |
| def encode_item(self, item: dict[str, Any]) -> EncodedWardrobeItem: | |
| metadata_text = self._build_item_prompt(item) | |
| cache_key = self._cache_key(item, metadata_text) | |
| cached = self._embedding_cache.get(cache_key) | |
| if cached is not None: | |
| self._embedding_cache.move_to_end(cache_key) | |
| return EncodedWardrobeItem( | |
| item=item, | |
| vector=cached.copy(), | |
| slot=infer_slot_name(item), | |
| metadata_text=metadata_text, | |
| ) | |
| text_vec = self.encode_text(metadata_text) | |
| image_vec = self.encode_image_url(str(item.get("image_url") or ""), metadata_text) | |
| vector = text_vec if image_vec is None else self._normalize_vector((text_vec + image_vec) / 2.0) | |
| self._remember_embedding(cache_key, vector) | |
| return EncodedWardrobeItem( | |
| item=item, | |
| vector=vector.copy(), | |
| slot=infer_slot_name(item), | |
| metadata_text=metadata_text, | |
| ) | |
| def encode_text(self, text: str) -> np.ndarray: | |
| self._ensure_model_loaded() | |
| if self._backend == "open_clip" and self._model is not None and self._tokenizer is not None: | |
| with torch.inference_mode(): | |
| tokens = self._tokenizer([text]).to(self.device) | |
| features = self._model.encode_text(tokens) | |
| return self._resize_and_normalize(features[0].detach().float().cpu().numpy()) | |
| if self._backend == "transformers" and self._model is not None and self._processor is not None: | |
| with torch.inference_mode(): | |
| inputs = self._processor( | |
| text=[text], | |
| return_tensors="pt", | |
| padding=True, | |
| truncation=True, | |
| ) | |
| inputs = {name: value.to(self.device) for name, value in inputs.items()} | |
| if hasattr(self._model, "get_text_features"): | |
| features = self._model.get_text_features(**inputs) | |
| else: | |
| outputs = self._model(**inputs) | |
| features = getattr(outputs, "pooler_output", outputs.last_hidden_state[:, 0, :]) | |
| return self._resize_and_normalize(features[0].detach().float().cpu().numpy()) | |
| return self._fallback_embedding(text) | |
| def encode_image_url(self, image_url: str, fallback_text: str) -> np.ndarray | None: | |
| image = self._load_image(image_url) | |
| if image is None: | |
| return None | |
| return self.encode_image(image, fallback_text) | |
| def encode_image(self, image: Image.Image, fallback_text: str = "") -> np.ndarray: | |
| self._ensure_model_loaded() | |
| if self._backend == "open_clip" and self._model is not None and self._preprocess is not None: | |
| with torch.inference_mode(): | |
| tensor = self._preprocess(image).unsqueeze(0).to(self.device) | |
| features = self._model.encode_image(tensor) | |
| return self._resize_and_normalize(features[0].detach().float().cpu().numpy()) | |
| if self._backend == "transformers" and self._model is not None and self._processor is not None: | |
| with torch.inference_mode(): | |
| inputs = self._processor(images=[image], return_tensors="pt") | |
| inputs = {name: value.to(self.device) for name, value in inputs.items()} | |
| if hasattr(self._model, "get_image_features"): | |
| features = self._model.get_image_features(**inputs) | |
| else: | |
| outputs = self._model(**inputs) | |
| features = getattr(outputs, "pooler_output", outputs.last_hidden_state[:, 0, :]) | |
| return self._resize_and_normalize(features[0].detach().float().cpu().numpy()) | |
| return self._fallback_embedding(f"image::{fallback_text}") | |
| def encode_context(self, context: RecommendationContext) -> np.ndarray: | |
| profile = context.user_profile or {} | |
| favorite_colors = profile.get("favorite_colors") | |
| disliked_styles = profile.get("disliked_styles") | |
| prompt_parts = [ | |
| f"An outfit for {context.occasion or 'casual'}", | |
| f"in {context.weather.season or 'all-season'} weather", | |
| f"temperature {context.weather.temperature_c}C" | |
| if context.weather.temperature_c is not None | |
| else "", | |
| "rainy conditions" if context.weather.is_rainy else "", | |
| f"region {context.region or 'global'}", | |
| f"preferred style {profile.get('style_profile')}" if profile.get("style_profile") else "", | |
| f"favorite colors {', '.join(favorite_colors)}" | |
| if isinstance(favorite_colors, list) and favorite_colors | |
| else "", | |
| f"disliked styles {', '.join(disliked_styles)}" | |
| if isinstance(disliked_styles, list) and disliked_styles | |
| else "", | |
| ] | |
| return self.encode_text(" ".join(part for part in prompt_parts if part)) | |
| def _ensure_model_loaded(self) -> None: | |
| if self._load_attempted: | |
| return | |
| self._load_attempted = True | |
| if open_clip is not None and self.model_id.lower().startswith("marqo/"): | |
| try: | |
| self._model, _, self._preprocess = open_clip.create_model_and_transforms( | |
| f"hf-hub:{self.model_id}", | |
| device=self.device, | |
| ) | |
| self._tokenizer = open_clip.get_tokenizer(f"hf-hub:{self.model_id}") | |
| self._model.eval() | |
| self._backend = "open_clip" | |
| return | |
| except Exception: | |
| self._model = None | |
| self._preprocess = None | |
| self._tokenizer = None | |
| try: | |
| self._processor = AutoProcessor.from_pretrained(self.model_id) | |
| self._model = AutoModel.from_pretrained(self.model_id).to(self.device).eval() | |
| self._backend = "transformers" | |
| except Exception: | |
| self._model = None | |
| self._processor = None | |
| self._backend = "fallback-hash" | |
| def _load_image(self, image_url: str) -> Image.Image | None: | |
| if not image_url or image_url.startswith("memory://") or image_url.startswith("data:"): | |
| return None | |
| parsed = urlparse(image_url) | |
| try: | |
| if parsed.scheme in {"http", "https"}: | |
| request = Request( | |
| image_url, | |
| headers={"User-Agent": "Mozilla/5.0", "Accept": "image/*,*/*;q=0.8"}, | |
| ) | |
| with urlopen(request, timeout=self.image_timeout_seconds) as response: | |
| raw = response.read() | |
| return Image.open(io.BytesIO(raw)).convert("RGB") | |
| if os.path.isfile(image_url): | |
| return Image.open(image_url).convert("RGB") | |
| except Exception: | |
| return None | |
| return None | |
| def _build_item_prompt(self, item: dict[str, Any]) -> str: | |
| description = item.get("description") if isinstance(item.get("description"), dict) else {} | |
| category = item.get("category") or description.get("category") or description.get("type") or "garment" | |
| color = item.get("color") or description.get("color") or "unknown color" | |
| pattern = item.get("pattern") or description.get("pattern") or "solid" | |
| fabric = item.get("fabric") or description.get("fabric") or "unknown fabric" | |
| fit = item.get("fit") or description.get("fit") or "regular" | |
| season = item.get("season") or description.get("season") or "all-season" | |
| style = item.get("style") or description.get("occasion") or description.get("style") or "casual" | |
| slot = infer_slot_name(item) | |
| return ( | |
| f"Fashion product photo of a {color} {pattern} {fabric} {category}, " | |
| f"{fit} fit, {style} style, suitable for {season}, worn as {slot}." | |
| ) | |
| def _cache_key(self, item: dict[str, Any], metadata_text: str) -> str: | |
| payload = { | |
| "id": str(item.get("id") or ""), | |
| "image_url": str(item.get("image_url") or ""), | |
| "metadata_text": metadata_text, | |
| } | |
| raw = json.dumps(payload, sort_keys=True, ensure_ascii=True).encode("utf-8") | |
| return hashlib.sha256(raw).hexdigest() | |
| def _remember_embedding(self, cache_key: str, vector: np.ndarray) -> None: | |
| self._embedding_cache[cache_key] = vector.copy() | |
| self._embedding_cache.move_to_end(cache_key) | |
| while len(self._embedding_cache) > self.cache_size: | |
| self._embedding_cache.popitem(last=False) | |
| def _fallback_embedding(self, seed_text: str) -> np.ndarray: | |
| digest = hashlib.sha256(seed_text.encode("utf-8", errors="ignore")).digest() | |
| seed = int.from_bytes(digest[:8], "big", signed=False) | |
| rng = np.random.default_rng(seed) | |
| return self._normalize_vector(rng.standard_normal(self.embedding_dim).astype(np.float32)) | |
| def _resize_and_normalize(self, vector: np.ndarray) -> np.ndarray: | |
| arr = np.asarray(vector, dtype=np.float32).reshape(-1) | |
| if arr.shape[0] == self.embedding_dim: | |
| return self._normalize_vector(arr) | |
| if arr.shape[0] < 2: | |
| return self._fallback_embedding(str(arr.tolist())) | |
| src_x = np.linspace(0.0, 1.0, num=arr.shape[0], dtype=np.float32) | |
| dst_x = np.linspace(0.0, 1.0, num=self.embedding_dim, dtype=np.float32) | |
| resized = np.interp(dst_x, src_x, arr).astype(np.float32) | |
| return self._normalize_vector(resized) | |
| def _normalize_vector(vector: np.ndarray) -> np.ndarray: | |
| arr = np.asarray(vector, dtype=np.float32).reshape(-1) | |
| norm = float(np.linalg.norm(arr)) | |
| if norm < 1e-8: | |
| return np.zeros_like(arr, dtype=np.float32) | |
| return arr / norm | |