from __future__ import annotations import json import math import os import random from pathlib import Path from typing import Any import numpy as np import torch import torch.nn.functional as F DATASET_NAME = "nlphuji/flickr30k" DATASET_SPLIT = "test" def seed_everything(seed: int) -> None: random.seed(seed) np.random.seed(seed) torch.manual_seed(seed) if torch.cuda.is_available(): torch.cuda.manual_seed_all(seed) def read_json(path: str | os.PathLike[str]) -> dict[str, Any]: with open(path, encoding="utf-8") as f: return json.load(f) def write_json(path: str | os.PathLike[str], value: Any) -> None: path = Path(path) path.parent.mkdir(parents=True, exist_ok=True) with open(path, "w", encoding="utf-8") as f: json.dump(value, f, indent=2, ensure_ascii=False) def normalized(x: torch.Tensor, eps: float = 1e-8) -> torch.Tensor: return F.normalize(x.float(), dim=-1, eps=eps) def pairwise_cosine_distance(x: np.ndarray) -> np.ndarray: x = x.astype(np.float64, copy=False) x /= np.linalg.norm(x, axis=1, keepdims=True).clip(min=1e-12) d = 1.0 - x @ x.T np.fill_diagonal(d, 0.0) scale = np.median(d[d > 0]) return d / max(float(scale), 1e-12) def linear_cka(x: torch.Tensor, y: torch.Tensor) -> float: x = x.float() - x.float().mean(0, keepdim=True) y = y.float() - y.float().mean(0, keepdim=True) xty = x.T @ y numerator = (xty * xty).sum() xx = x.T @ x yy = y.T @ y denominator = torch.sqrt((xx * xx).sum() * (yy * yy).sum()) return float((numerator / denominator.clamp_min(1e-12)).item()) def retrieval_metrics( image_features: torch.Tensor, text_features: torch.Tensor, ks: tuple[int, ...] = (1, 5, 10), ) -> dict[str, float]: image_features = normalized(image_features) text_features = normalized(text_features) similarities = image_features @ text_features.T n = similarities.shape[0] truth = torch.arange(n, device=similarities.device) i2t_order = similarities.argsort(dim=1, descending=True) t2i_order = similarities.T.argsort(dim=1, descending=True) i2t_rank = (i2t_order == truth[:, None]).nonzero()[:, 1] t2i_rank = (t2i_order == truth[:, None]).nonzero()[:, 1] result: dict[str, float] = {} for k in ks: result[f"i2t_r@{k}"] = float((i2t_rank < k).float().mean().item()) result[f"t2i_r@{k}"] = float((t2i_rank < k).float().mean().item()) result["i2t_median_rank"] = float(i2t_rank.float().median().item() + 1) result["t2i_median_rank"] = float(t2i_rank.float().median().item() + 1) result["chance_r@1"] = 1.0 / max(n, 1) return result def batch_indices(n: int, batch_size: int, shuffle: bool = False, seed: int = 0): order = np.arange(n) if shuffle: rng = np.random.default_rng(seed) rng.shuffle(order) for start in range(0, n, batch_size): yield order[start : start + batch_size] def sliced_wasserstein( x: torch.Tensor, y: torch.Tensor, num_projections: int = 64, ) -> torch.Tensor: """Differentiable empirical sliced W2 for equal-sized minibatches.""" n = min(x.shape[0], y.shape[0]) x = x[:n] y = y[:n] directions = torch.randn( x.shape[-1], num_projections, device=x.device, dtype=x.dtype ) directions = F.normalize(directions, dim=0) x_proj = (x @ directions).sort(dim=0).values y_proj = (y @ directions).sort(dim=0).values return (x_proj - y_proj).square().mean() def cosine_isometry_loss(source: torch.Tensor, mapped: torch.Tensor) -> torch.Tensor: source = normalized(source) mapped = normalized(mapped) source_gram = source @ source.T mapped_gram = mapped @ mapped.T mask = ~torch.eye(source.shape[0], dtype=torch.bool, device=source.device) return (source_gram[mask] - mapped_gram[mask]).square().mean() def cosine_loss(x: torch.Tensor, y: torch.Tensor) -> torch.Tensor: return 1.0 - F.cosine_similarity(x.float(), y.float(), dim=-1).mean() def dtype_for_device(device: str) -> torch.dtype: return torch.bfloat16 if device.startswith("cuda") else torch.float32 def parameter_count(module: torch.nn.Module) -> int: return sum(p.numel() for p in module.parameters()) def cosine_schedule(step: int, steps: int, warmup: int) -> float: if step < warmup: return (step + 1) / max(1, warmup) progress = (step - warmup) / max(1, steps - warmup) return 0.5 * (1.0 + math.cos(math.pi * progress))