Spaces:
Sleeping
Sleeping
| from __future__ import annotations | |
| from typing import Any | |
| import numpy as np | |
| from .encoder import FashionItemEncoder | |
| from .schemas import EncodedWardrobeItem, RecommendationContext, SlotName | |
| class OutfitCandidateRetriever: | |
| """Slot-aware embedding retrieval with MMR diversification.""" | |
| def __init__( | |
| self, | |
| encoder: FashionItemEncoder, | |
| slot_pool_size: int = 24, | |
| mmr_lambda: float = 0.72, | |
| ) -> None: | |
| self.encoder = encoder | |
| self.slot_pool_size = slot_pool_size | |
| self.mmr_lambda = mmr_lambda | |
| def encode_wardrobe(self, wardrobe_items: list[dict[str, Any]]) -> list[EncodedWardrobeItem]: | |
| return [self.encoder.encode_item(item) for item in wardrobe_items] | |
| def split_by_slot( | |
| self, | |
| encoded_items: list[EncodedWardrobeItem], | |
| ) -> dict[SlotName, list[EncodedWardrobeItem]]: | |
| buckets: dict[SlotName, list[EncodedWardrobeItem]] = { | |
| "top": [], | |
| "bottom": [], | |
| "shoes": [], | |
| "accessory": [], | |
| "unknown": [], | |
| } | |
| for item in encoded_items: | |
| buckets[item.slot].append(item) | |
| return buckets | |
| def retrieve( | |
| self, | |
| encoded_items: list[EncodedWardrobeItem], | |
| context: RecommendationContext, | |
| locked_top: dict[str, Any] | None = None, | |
| locked_bottom: dict[str, Any] | None = None, | |
| locked_other: dict[str, Any] | None = None, | |
| candidate_pool: int | None = None, | |
| ) -> dict[SlotName, list[EncodedWardrobeItem]]: | |
| buckets = self.split_by_slot(encoded_items) | |
| query_vec = self.encoder.encode_context(context) | |
| pool_size = candidate_pool or self.slot_pool_size | |
| locked_top_vec = self.encoder.encode_item(locked_top).vector if locked_top else None | |
| locked_bottom_vec = self.encoder.encode_item(locked_bottom).vector if locked_bottom else None | |
| locked_other_encoded = self.encoder.encode_item(locked_other) if locked_other else None | |
| accessory_bucket = buckets["accessory"] + buckets["unknown"] | |
| return { | |
| "top": [self.encoder.encode_item(locked_top)] if locked_top else self._rank_bucket( | |
| buckets["top"], | |
| query_vec=self._merge_query(query_vec, locked_bottom_vec), | |
| top_k=pool_size, | |
| ), | |
| "bottom": [self.encoder.encode_item(locked_bottom)] if locked_bottom else self._rank_bucket( | |
| buckets["bottom"], | |
| query_vec=self._merge_query(query_vec, locked_top_vec), | |
| top_k=pool_size, | |
| ), | |
| "shoes": self._rank_bucket( | |
| buckets["shoes"], | |
| query_vec=self._merge_query( | |
| query_vec, | |
| locked_top_vec, | |
| locked_bottom_vec, | |
| locked_other_encoded.vector if locked_other_encoded and locked_other_encoded.slot == "shoes" else None, | |
| ), | |
| top_k=min(pool_size, 12), | |
| ) if not locked_other_encoded or locked_other_encoded.slot != "shoes" else [locked_other_encoded], | |
| "accessory": ( | |
| [locked_other_encoded] | |
| if locked_other_encoded and locked_other_encoded.slot != "shoes" | |
| else self._rank_bucket( | |
| accessory_bucket, | |
| query_vec=self._merge_query(query_vec, locked_top_vec, locked_bottom_vec), | |
| top_k=min(pool_size, 12), | |
| ) | |
| ), | |
| "unknown": [], | |
| } | |
| def _rank_bucket( | |
| self, | |
| bucket: list[EncodedWardrobeItem], | |
| query_vec: np.ndarray, | |
| top_k: int, | |
| ) -> list[EncodedWardrobeItem]: | |
| if not bucket: | |
| return [] | |
| ranked = sorted( | |
| bucket, | |
| key=lambda item: float(np.dot(item.vector, query_vec)), | |
| reverse=True, | |
| ) | |
| return self._mmr_diversify(ranked, query_vec=query_vec, top_k=top_k) | |
| def _mmr_diversify( | |
| self, | |
| ranked_items: list[EncodedWardrobeItem], | |
| query_vec: np.ndarray, | |
| top_k: int, | |
| ) -> list[EncodedWardrobeItem]: | |
| selected: list[EncodedWardrobeItem] = [] | |
| remaining = list(ranked_items) | |
| while remaining and len(selected) < top_k: | |
| best_index = 0 | |
| best_score = -1e9 | |
| for index, item in enumerate(remaining): | |
| query_score = float(np.dot(item.vector, query_vec)) | |
| novelty_penalty = 0.0 | |
| if selected: | |
| novelty_penalty = max(float(np.dot(item.vector, prev.vector)) for prev in selected) | |
| mmr_score = self.mmr_lambda * query_score - (1.0 - self.mmr_lambda) * novelty_penalty | |
| if mmr_score > best_score: | |
| best_index = index | |
| best_score = mmr_score | |
| selected.append(remaining.pop(best_index)) | |
| return selected | |
| def _merge_query(*vectors: np.ndarray | None) -> np.ndarray: | |
| valid = [np.asarray(vec, dtype=np.float32) for vec in vectors if vec is not None] | |
| if not valid: | |
| return np.zeros((1,), dtype=np.float32) | |
| merged = np.mean(np.stack(valid, axis=0), axis=0) | |
| norm = float(np.linalg.norm(merged)) | |
| if norm < 1e-8: | |
| return merged | |
| return merged / norm | |