StyleWellBackend / fashion_ai /retriever.py
HelloWorld0204's picture
Upload 16 files
e08551d verified
Raw
History Blame Contribute Delete
5.55 kB
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