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 | """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)) | |
| 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() | |
| 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) | |
| 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)]) | |