AbstractPhil's picture
vit-captionbank ship: two-stage cores (reusable) + bank expansions + cotrained pairs + frames + fidelity-gated loader + complete ledgers incl. the preregistered probe refutation with full controls
f78e030 verified
Raw
History Blame Contribute Delete
6.79 kB
"""Standalone loader for geolip-vit-captionbank-coco (no repo imports).
Two-stage design: the CORE student is a reusable 8.66M ViT encoder into the
5-CLIP GPA consensus space; the ALIGNMENT BANK is a modular EXPANSION that
appends a 128-d geometric context (640-d enriched output).
from loader import load_student, load_bank, embed_images, enrich
student = load_student("core/student_s0.pt")
emb = embed_images(student, images01) # (B, 512) unit vectors
bank = load_bank("banks/bank_s0.pt")
enriched = enrich(bank, emb) # (B, 640)
"""
import torch
import torch.nn as nn
import torch.nn.functional as F
CLIP_MEAN = (0.48145466, 0.4578275, 0.40821073)
CLIP_STD = (0.26862954, 0.26130258, 0.27577711)
class Block(nn.Module):
def __init__(self, d, heads):
super().__init__()
self.n1 = nn.LayerNorm(d)
self.qkv = nn.Linear(d, 3 * d)
self.proj = nn.Linear(d, d)
self.n2 = nn.LayerNorm(d)
self.fc1 = nn.Linear(d, 4 * d)
self.fc2 = nn.Linear(4 * d, d)
self.heads = heads
def forward(self, x):
B, N, C = x.shape
q, k, v = (self.qkv(self.n1(x))
.reshape(B, N, 3, self.heads, C // self.heads)
.permute(2, 0, 3, 1, 4))
a = F.scaled_dot_product_attention(q, k, v)
x = x + self.proj(a.transpose(1, 2).reshape(B, N, C))
return x + self.fc2(F.gelu(self.fc1(self.n2(x))))
class Student(nn.Module):
"""CLS-token readout ViT-Ti, linear head into the 512-d consensus."""
def __init__(self, out_dim=512, d=240, depth=12, heads=4, patch=16,
img=160):
super().__init__()
self.patch = nn.Conv2d(3, d, patch, patch)
self.cls = nn.Parameter(torch.zeros(1, 1, d))
self.pos = nn.Parameter(torch.zeros(1, (img // patch) ** 2 + 1, d))
self.blocks = nn.ModuleList(Block(d, heads) for _ in range(depth))
self.norm = nn.LayerNorm(d)
self.head = nn.Linear(d, out_dim)
def forward(self, x):
x = self.patch(x).flatten(2).transpose(1, 2)
x = torch.cat([self.cls.expand(x.shape[0], -1, -1), x], 1) + self.pos
for b in self.blocks:
x = b(x)
return self.head(self.norm(x)[:, 0])
class AlignmentBank(nn.Module):
"""The expansion: 5 whitened-Procrustes expert frames + 512 dense-cosine
anchors -> 538-d geometric signature -> 128-d context appended to the
embedding. Forward-only port (no training losses)."""
def __init__(self, d_embed=512, n_experts=5, n_anchors=512, d_bank=128):
super().__init__()
self.d_embed, self.n_experts = d_embed, n_experts
self.expert_rotations = nn.ParameterList(
[nn.Parameter(torch.eye(d_embed)) for _ in range(n_experts)])
self.expert_whiteners = nn.ParameterList(
[nn.Parameter(torch.eye(d_embed)) for _ in range(n_experts)])
self.expert_means = nn.ParameterList(
[nn.Parameter(torch.zeros(d_embed)) for _ in range(n_experts)])
self.anchors = nn.Parameter(
F.normalize(torch.randn(n_anchors, d_embed), dim=-1))
geo_dim = n_experts * 3 + n_experts * (n_experts - 1) // 2 + 1 \
+ n_anchors
self.geo_proj = nn.Sequential(
nn.Linear(geo_dim, d_bank * 2), nn.GELU(),
nn.LayerNorm(d_bank * 2),
nn.Linear(d_bank * 2, d_bank), nn.LayerNorm(d_bank))
self.register_buffer("target_cv", torch.tensor(0.20))
self.register_buffer("target_cross_cos_mean", torch.tensor(0.0))
self.register_buffer("target_cross_cos_std", torch.tensor(0.0))
self.register_buffer("target_disagreement_ratio", torch.tensor(0.0))
@torch.no_grad()
def forward(self, embedding):
emb = embedding.float()
cons, recon, proj, norms = [], [], [], []
for i in range(self.n_experts):
R, W, mu = (self.expert_rotations[i], self.expert_whiteners[i],
self.expert_means[i])
whitened = (emb - mu) @ W
wn = F.normalize(whitened, dim=-1)
in_expert = wn @ R.T
back = in_expert @ R
cons.append(F.cosine_similarity(wn, back, dim=-1))
recon.append((wn - back).pow(2).mean(dim=-1))
proj.append(in_expert)
norms.append(whitened.norm(dim=-1))
expert_cos = torch.stack(cons, -1)
expert_mse = torch.stack(recon, -1)
cross = torch.stack(
[F.cosine_similarity(proj[i], proj[j], dim=-1)
for i in range(self.n_experts)
for j in range(i + 1, self.n_experts)], -1)
ratio = expert_cos.std(-1) / (expert_cos.mean(-1) + 1e-8)
norm_ratio = norms_t = torch.stack(norms, -1)
norm_ratio = norms_t / (norms_t.mean(-1, keepdim=True) + 1e-8)
anchor_cos = emb @ F.normalize(self.anchors, dim=-1).T
sig = torch.cat([expert_cos, expert_mse, cross,
ratio.unsqueeze(-1), norm_ratio, anchor_cos], -1)
return torch.cat([embedding, self.geo_proj(sig)], -1)
def load_student(path, device=None):
device = device or ("cuda" if torch.cuda.is_available() else "cpu")
ck = torch.load(path, map_location="cpu", weights_only=True)
sd = ck["state_dict"] if "state_dict" in ck else ck
m = Student(out_dim=sd["head.weight"].shape[0])
m.load_state_dict(sd, strict=True)
return m.to(device).eval()
def load_bank(path, device=None):
device = device or ("cuda" if torch.cuda.is_available() else "cpu")
ck = torch.load(path, map_location="cpu", weights_only=True)
b = AlignmentBank()
b.load_state_dict(ck["state_dict"], strict=True)
return b.to(device).eval()
@torch.no_grad()
def embed_images(model, images, batch=256):
"""images: list of PIL.Image (any size). Returns L2-normalized (B,512).
Preprocessing is the training pipeline verbatim: shortest-side
Resize(182, bicubic) -> CenterCrop(160) -> CLIP normalization."""
from torchvision import transforms as T
from torchvision.transforms import InterpolationMode
tf = T.Compose([
T.Resize(182, interpolation=InterpolationMode.BICUBIC),
T.CenterCrop(160), T.ToTensor(),
T.Normalize(CLIP_MEAN, CLIP_STD)])
dev = next(model.parameters()).device
out = []
for i in range(0, len(images), batch):
x = torch.stack([tf(im.convert("RGB"))
for im in images[i:i + batch]]).to(dev)
out.append(F.normalize(model(x), dim=-1).cpu())
return torch.cat(out)
@torch.no_grad()
def enrich(bank, emb, batch=4096):
dev = next(bank.parameters()).device
return torch.cat([bank(emb[i:i + batch].to(dev)).cpu()
for i in range(0, len(emb), batch)])