- 移除 JEPA/lejepa-identifiability 子模块 gitlink - 移除 research/multiply/MultiPLY 子模块 gitlink - 删除 .gitmodules(不再有外部 URL 依赖) - 两个目录内容作为普通文件纳入主仓库追踪 - 删除各自内部 .git 目录,消除嵌套 git 仓库
This commit is contained in:
@@ -0,0 +1,53 @@
|
||||
"""Loss functions."""
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
|
||||
class SIGReg(nn.Module):
|
||||
"""Sliced characteristic function regularizer (Balestriero & LeCun 2025)."""
|
||||
|
||||
def __init__(self, knots=17, n_slices=256, t_max=3.0):
|
||||
super().__init__()
|
||||
self.n_slices = n_slices
|
||||
t = torch.linspace(0, t_max, knots)
|
||||
dt = t_max / (knots - 1)
|
||||
w = torch.full((knots,), 2 * dt)
|
||||
w[[0, -1]] = dt
|
||||
self.register_buffer("t", t)
|
||||
self.register_buffer("phi", torch.exp(-t**2 / 2))
|
||||
self.register_buffer("weights", w * torch.exp(-t**2 / 2))
|
||||
|
||||
def forward(self, h):
|
||||
"""h: (V, B, N) -> scalar."""
|
||||
flat = h.flatten(0, 1)
|
||||
A = F.normalize(torch.randn(flat.size(-1), self.n_slices, device=flat.device), dim=0)
|
||||
xt = (flat @ A).unsqueeze(-1) * self.t
|
||||
err = (xt.cos().mean(0) - self.phi) ** 2 + xt.sin().mean(0) ** 2
|
||||
return (err @ self.weights).mean() * flat.size(0)
|
||||
|
||||
|
||||
def whitening_loss(h):
|
||||
"""||Cov(h) - I||²_F. h: (V, B, N) -> scalar."""
|
||||
flat = h.flatten(0, 1)
|
||||
flat = flat - flat.mean(dim=0)
|
||||
cov = (flat.T @ flat) / (flat.shape[0] - 1)
|
||||
return (cov - torch.eye(flat.shape[1], device=h.device)).square().mean()
|
||||
|
||||
|
||||
def alignment_loss(h):
|
||||
"""Pull positive-pair views together. h: (V, B, N) -> scalar."""
|
||||
return (h.mean(0) - h).square().mean()
|
||||
|
||||
|
||||
def infonce_loss(h, sigma):
|
||||
"""Symmetric Gaussian-kernel InfoNCE: sim(u, v) = -||u - v||² / (2σ²).
|
||||
h: (V, B, N) with V=2 views. Negatives are other batch elements.
|
||||
"""
|
||||
h1, h2 = h[0], h[1] # (B, N) each
|
||||
d12 = ((h1.unsqueeze(1) - h2.unsqueeze(0)) ** 2).sum(-1) # (B, B)
|
||||
sim = -d12 / (2 * sigma ** 2)
|
||||
loss_a = -(sim.diag() - torch.logsumexp(sim, dim=1)).mean()
|
||||
loss_b = -(sim.diag() - torch.logsumexp(sim, dim=0)).mean()
|
||||
return 0.5 * (loss_a + loss_b)
|
||||
Reference in New Issue
Block a user