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 @staticmethod 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