CharlesCNorton
Image-level person classification on EUPE-ViT-B features with no free parameters
e8b8483
Raw
History Blame Contribute Delete
1.74 kB
"""Pooled feature extraction and the comparison rule.
One 768-D vector per image: layernorm across the 768 channels of every patch
token, then max-pool across patches. The decision compares two sums of that
vector against each other, so the boundary sits at zero and carries no offset.
"""
from typing import Sequence
import torch
import torch.nn.functional as F
D = 768
RES = 768
def pool(patch_tokens: torch.Tensor) -> torch.Tensor:
"""(N, D) or (B, N, D) patch tokens -> (D,) or (B, D) pooled vector."""
ln = F.layer_norm(patch_tokens.float(), [D])
return ln.max(dim=-2).values
@torch.inference_mode()
def backbone_pooled(backbone, x: torch.Tensor, autocast: bool = True) -> torch.Tensor:
"""Forward a normalized batch through the backbone and pool it."""
if autocast:
dev = 'cuda' if x.is_cuda else 'cpu'
with torch.autocast(dev, dtype=torch.bfloat16):
out = backbone.forward_features(x)
else:
out = backbone.forward_features(x)
return pool(out['x_norm_patchtokens'].float())
def _as_index(idx, like: torch.Tensor) -> torch.Tensor:
if torch.is_tensor(idx):
return idx
return torch.tensor(list(idx), dtype=torch.long, device=like.device)
def score(pooled: torch.Tensor, pos: Sequence[int], neg: Sequence[int]) -> torch.Tensor:
"""sum(pooled[pos]) - sum(pooled[neg]), over the last axis."""
p, n = _as_index(pos, pooled), _as_index(neg, pooled)
return pooled.index_select(-1, p).sum(-1) - pooled.index_select(-1, n).sum(-1)
def decide(pooled: torch.Tensor, pos: Sequence[int], neg: Sequence[int]) -> torch.Tensor:
"""sum(pooled[pos]) > sum(pooled[neg]). No threshold, no free parameter."""
return score(pooled, pos, neg) > 0