Spaces:
Sleeping
Sleeping
| from __future__ import annotations | |
| import json | |
| import random | |
| from dataclasses import dataclass | |
| from pathlib import Path | |
| from typing import Any, Iterable | |
| import numpy as np | |
| import torch | |
| import torch.nn.functional as F | |
| from torch.utils.data import DataLoader, Dataset | |
| from .encoder import FashionItemEncoder, infer_slot_name | |
| from .ranker import OutfitCompatibilityRanker | |
| from .schemas import RecommendationContext, WeatherContext | |
| class TrainingSample: | |
| outfit: dict[str, Any] | |
| label: float | |
| occasion: str = "casual" | |
| weather: dict[str, Any] | None = None | |
| user_profile: dict[str, Any] | None = None | |
| class OutfitRankingDataset(Dataset[TrainingSample]): | |
| """ | |
| JSONL schema per row: | |
| { | |
| "outfit": {"top": {...}, "bottom": {...}, "shoes": {...}, "accessory": {...}}, | |
| "label": 1, | |
| "occasion": "formal", | |
| "weather": {"season": "summer", "temperature_c": 30}, | |
| "user_profile": {"style_profile": "minimal", "favorite_colors": ["navy"]} | |
| } | |
| """ | |
| def __init__(self, path: str | Path) -> None: | |
| self.samples = self._load_samples(path) | |
| def __len__(self) -> int: | |
| return len(self.samples) | |
| def __getitem__(self, index: int) -> TrainingSample: | |
| return self.samples[index] | |
| def _load_samples(path: str | Path) -> list[TrainingSample]: | |
| rows: list[TrainingSample] = [] | |
| with open(path, "r", encoding="utf-8") as file_obj: | |
| for line in file_obj: | |
| line = line.strip() | |
| if not line: | |
| continue | |
| payload = json.loads(line) | |
| rows.append( | |
| TrainingSample( | |
| outfit=payload.get("outfit", {}), | |
| label=float(payload.get("label", 1.0)), | |
| occasion=str(payload.get("occasion") or "casual"), | |
| weather=payload.get("weather") if isinstance(payload.get("weather"), dict) else None, | |
| user_profile=payload.get("user_profile") | |
| if isinstance(payload.get("user_profile"), dict) | |
| else None, | |
| ) | |
| ) | |
| return rows | |
| class NegativeSampler: | |
| """Hard negative sampler using nearest same-slot replacements.""" | |
| def __init__( | |
| self, | |
| catalog_items: list[dict[str, Any]], | |
| encoder: FashionItemEncoder, | |
| hard_negative_top_k: int = 20, | |
| ) -> None: | |
| self.encoder = encoder | |
| self.hard_negative_top_k = hard_negative_top_k | |
| self.slot_catalog = { | |
| "top": [], | |
| "bottom": [], | |
| "shoes": [], | |
| "accessory": [], | |
| "unknown": [], | |
| } | |
| for item in catalog_items: | |
| encoded = self.encoder.encode_item(item) | |
| self.slot_catalog[encoded.slot].append(encoded) | |
| def sample(self, outfit: dict[str, Any], occasion: str = "casual") -> dict[str, Any]: | |
| candidate = dict(outfit) | |
| replaceable_slots = [ | |
| slot_name | |
| for slot_name in ["top", "bottom", "shoes", "accessory"] | |
| if isinstance(outfit.get(slot_name), dict) | |
| ] | |
| if not replaceable_slots: | |
| return candidate | |
| slot_to_replace = random.choice(replaceable_slots) | |
| anchor = outfit[slot_to_replace] | |
| pool = self.slot_catalog.get(infer_slot_name(anchor), []) | |
| if not pool: | |
| return candidate | |
| anchor_vec = self.encoder.encode_item(anchor).vector | |
| ranked = sorted(pool, key=lambda entry: float(np.dot(entry.vector, anchor_vec)), reverse=True) | |
| hard_pool = [ | |
| entry.item | |
| for entry in ranked[: self.hard_negative_top_k] | |
| if str(entry.item.get("id")) != str(anchor.get("id")) | |
| ] | |
| if not hard_pool: | |
| return candidate | |
| candidate[slot_to_replace] = random.choice(hard_pool) | |
| candidate["negative_source"] = { | |
| "strategy": "hard_same_slot_replacement", | |
| "slot": slot_to_replace, | |
| "occasion": occasion, | |
| } | |
| return candidate | |
| class OutfitRankerCollator: | |
| def __init__(self, encoder: FashionItemEncoder) -> None: | |
| self.encoder = encoder | |
| def __call__(self, samples: Iterable[TrainingSample]) -> dict[str, torch.Tensor]: | |
| token_rows = [] | |
| mask_rows = [] | |
| labels = [] | |
| for sample in samples: | |
| context = RecommendationContext( | |
| occasion=sample.occasion, | |
| weather=WeatherContext( | |
| season=str((sample.weather or {}).get("season") or "all-season"), | |
| temperature_c=(sample.weather or {}).get("temperature_c"), | |
| is_rainy=(sample.weather or {}).get("is_rainy"), | |
| ), | |
| user_profile=sample.user_profile or {}, | |
| ) | |
| row = [ | |
| self.encoder.encode_context(context), | |
| self.encoder.encode_text(json.dumps(sample.user_profile or {"style_profile": "general"}, sort_keys=True)), | |
| ] | |
| mask = [1, 1] | |
| for slot_name in ["top", "bottom", "shoes", "accessory"]: | |
| slot_value = sample.outfit.get(slot_name) | |
| if isinstance(slot_value, dict): | |
| row.append(self.encoder.encode_item(slot_value).vector) | |
| mask.append(1) | |
| else: | |
| row.append(np.zeros(self.encoder.embedding_dim, dtype=np.float32)) | |
| mask.append(0) | |
| token_rows.append(np.stack(row, axis=0)) | |
| mask_rows.append(mask) | |
| labels.append(sample.label) | |
| return { | |
| "outfit_tokens": torch.tensor(np.stack(token_rows), dtype=torch.float32), | |
| "attention_mask": torch.tensor(np.asarray(mask_rows), dtype=torch.long), | |
| "labels": torch.tensor(np.asarray(labels), dtype=torch.float32), | |
| } | |
| def bpr_pairwise_loss(pos_logits: torch.Tensor, neg_logits: torch.Tensor) -> torch.Tensor: | |
| return -F.logsigmoid(pos_logits - neg_logits).mean() | |
| def train_ranker( | |
| train_jsonl: str | Path, | |
| catalog_jsonl: str | Path, | |
| output_checkpoint: str | Path, | |
| encoder_model_id: str = "patrickjohncyh/fashion-clip", | |
| epochs: int = 5, | |
| batch_size: int = 16, | |
| lr: float = 2e-4, | |
| device: str | None = None, | |
| ) -> None: | |
| device = device or ("cuda" if torch.cuda.is_available() else "cpu") | |
| encoder = FashionItemEncoder(model_id=encoder_model_id, device=device) | |
| dataset = OutfitRankingDataset(train_jsonl) | |
| negative_sampler = NegativeSampler(_load_catalog(catalog_jsonl), encoder) | |
| collator = OutfitRankerCollator(encoder) | |
| model = OutfitCompatibilityRanker(d_model=encoder.embedding_dim).to(device) | |
| optimizer = torch.optim.AdamW(model.parameters(), lr=lr, weight_decay=0.01) | |
| for _epoch in range(epochs): | |
| model.train() | |
| augmented_samples = [] | |
| for sample in dataset.samples: | |
| augmented_samples.append(sample) | |
| augmented_samples.append( | |
| TrainingSample( | |
| outfit=negative_sampler.sample(sample.outfit, sample.occasion), | |
| label=0.0, | |
| occasion=sample.occasion, | |
| weather=sample.weather, | |
| user_profile=sample.user_profile, | |
| ) | |
| ) | |
| loader = DataLoader( | |
| augmented_samples, | |
| batch_size=batch_size, | |
| shuffle=True, | |
| collate_fn=collator, | |
| ) | |
| for batch in loader: | |
| logits = model( | |
| batch["outfit_tokens"].to(device), | |
| batch["attention_mask"].to(device), | |
| ).squeeze(-1) | |
| labels = batch["labels"].to(device) | |
| bce_loss = F.binary_cross_entropy_with_logits(logits, labels) | |
| pos_logits = logits[labels > 0.5] | |
| neg_logits = logits[labels <= 0.5] | |
| if len(pos_logits) > 0 and len(neg_logits) > 0: | |
| limit = min(len(pos_logits), len(neg_logits)) | |
| pairwise = bpr_pairwise_loss(pos_logits[:limit], neg_logits[:limit]) | |
| else: | |
| pairwise = torch.zeros((), device=device) | |
| loss = bce_loss + 0.4 * pairwise | |
| optimizer.zero_grad(set_to_none=True) | |
| loss.backward() | |
| torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0) | |
| optimizer.step() | |
| Path(output_checkpoint).parent.mkdir(parents=True, exist_ok=True) | |
| torch.save( | |
| { | |
| "model_state_dict": model.state_dict(), | |
| "encoder_model_id": encoder_model_id, | |
| "embedding_dim": encoder.embedding_dim, | |
| }, | |
| output_checkpoint, | |
| ) | |
| def _load_catalog(path: str | Path) -> list[dict[str, Any]]: | |
| records = [] | |
| with open(path, "r", encoding="utf-8") as file_obj: | |
| for line in file_obj: | |
| line = line.strip() | |
| if line: | |
| records.append(json.loads(line)) | |
| return records | |