summaryrefslogtreecommitdiff
path: root/worldalign/common.py
diff options
context:
space:
mode:
Diffstat (limited to 'worldalign/common.py')
-rw-r--r--worldalign/common.py142
1 files changed, 142 insertions, 0 deletions
diff --git a/worldalign/common.py b/worldalign/common.py
new file mode 100644
index 0000000..aa7b174
--- /dev/null
+++ b/worldalign/common.py
@@ -0,0 +1,142 @@
+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))
+