chipoint / run_chipoint.py
4latte's picture
Upload folder using huggingface_hub
90c4a87 verified
Raw
History Blame Contribute Delete
7.91 kB
#!/usr/bin/env python3
"""
Chipoint — run the image-geolocation model on your own photo(s).
Encodes each image with three frozen vision encoders, projects through trained heads, retrieves the nearest
gallery images, and applies the density-weighted consensus that produced the benchmark numbers. Prints a
predicted (latitude, longitude) per image.
Batch-friendly: pass many images and each 5 GB gallery index is read exactly ONCE for the whole batch (all
images are scored against it in a single pass), instead of re-reading 15 GB per image.
Usage:
python run_chipoint.py path/to/photo.jpg [more.jpg ...] # globs are expanded, e.g. photos/*.jpg
Weights (heads + gallery indexes + coords) must sit next to this script (as when you download the repo).
Override the location with CHIPOINT_DIR=/path/to/weights. DINOv3 is pulled ungated from ModelScope unless
DINOV3_PATH points at a local copy.
"""
import os, sys, glob
os.environ.setdefault("PYTORCH_HIP_ALLOC_CONF", "expandable_segments:True")
os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF", "expandable_segments:True")
from pathlib import Path
import numpy as np, torch, torch.nn as nn, torch.nn.functional as F
from PIL import Image
W = Path(os.environ.get("CHIPOINT_DIR") or Path(__file__).resolve().parent) # weights live next to this script
DEV = "cuda" if torch.cuda.is_available() else "cpu"
AC = torch.bfloat16 if DEV == "cuda" else torch.float32
M, K = 200, 64 # retrieve top-M per encoder, keep top-K fused candidates
Gv, LAM, MU, BW1, BW2 = 0.8, 1.7, 0.8, 1492.7, 750.0 # locked density-weighted-consensus config
ENC_D = {"so400m256": 1152, "dinov3": 1024, "gopt": 1536}
class Head(nn.Module):
def __init__(s, D):
super().__init__(); s.net = nn.Sequential(nn.Linear(D, 1024), nn.GELU(), nn.Dropout(0.1), nn.Linear(1024, 512))
def forward(s, x): return F.normalize(s.net(x), dim=-1)
def load_head(tag, D):
h = Head(D); sd = torch.load(W / f"proj_head_{tag}.pt", map_location="cpu")
h.load_state_dict({k: v for k, v in sd.items() if k.startswith("net.")}) # drop geocell centers, unused at inference
return h.to(DEV).eval()
# ---- encode ALL images with the 3 encoders, then release the encoders (free VRAM for retrieval) ----------
@torch.no_grad()
def encode_all(imgs, bs=64):
import open_clip, torchvision.transforms as T
from transformers import AutoModel, AutoImageProcessor
E = {}
for tag, name in [("so400m256", "ViT-SO400M-16-SigLIP2-256"), ("gopt", "ViT-gopt-16-SigLIP2-256")]:
m, _, pre = open_clip.create_model_and_transforms(name, pretrained="webli"); m = m.to(DEV).eval(); out = []
for i in range(0, len(imgs), bs):
x = torch.stack([pre(p) for p in imgs[i:i+bs]]).to(DEV)
with torch.autocast(DEV, dtype=AC): e = m.encode_image(x).float()
out.append(F.normalize(e, dim=-1).cpu().numpy().astype("float32"))
E[tag] = np.concatenate(out); del m
if DEV == "cuda": torch.cuda.empty_cache()
src = os.environ.get("DINOV3_PATH")
if not src:
from modelscope import snapshot_download as _ms
src = _ms("facebook/dinov3-vitl16-pretrain-lvd1689m")
dm = AutoModel.from_pretrained(src).to(DEV).eval()
if DEV == "cuda": dm = dm.half()
proc = AutoImageProcessor.from_pretrained(src)
mean = torch.tensor(proc.image_mean).view(1, 3, 1, 1); std = torch.tensor(proc.image_std).view(1, 3, 1, 1)
pre = T.Compose([T.Resize(256), T.CenterCrop(224), T.ToTensor()]); out = []
for i in range(0, len(imgs), bs):
x = torch.stack([(pre(p) - mean[0]) / std[0] for p in imgs[i:i+bs]]).to(DEV)
x = x.half() if DEV == "cuda" else x
o = dm(x); e = o.pooler_output if getattr(o, "pooler_output", None) is not None else o.last_hidden_state[:, 0]
out.append(F.normalize(e.float(), dim=-1).cpu().numpy().astype("float32"))
E["dinov3"] = np.concatenate(out); del dm
if DEV == "cuda": torch.cuda.empty_cache()
return E
@torch.no_grad()
def project(tag, emb):
return F.normalize(HEADS[tag](torch.tensor(emb).to(DEV)), dim=-1).cpu().numpy().astype("float32")
# ---- retrieve ALL queries against one gallery in a SINGLE pass (gallery read once) -----------------------
@torch.no_grad()
def retrieve_all(tag, Q): # Q:[N,512] -> idx[N,M], sim[N,M]
g = np.load(W / f"proj_head_gallery_512_{tag}.npy", mmap_mode="r"); N = len(Q)
if DEV == "cuda":
qt = torch.from_numpy(Q).to(DEV).half()
bs = torch.full((N, M), -1e9, device=DEV); bi = torch.zeros((N, M), dtype=torch.long, device=DEV); step = 250000
for i in range(0, g.shape[0], step):
gg = torch.from_numpy(np.asarray(g[i:i+step], dtype="float16")).to(DEV)
s = (qt @ gg.T).float(); cs, ci = s.topk(min(M, s.shape[1]), 1); ci = ci + i
al = torch.cat([bs, cs], 1); ai = torch.cat([bi, ci], 1); o = al.topk(M, 1).indices
bs = torch.gather(al, 1, o); bi = torch.gather(ai, 1, o); del gg, s
torch.cuda.empty_cache()
return bi.cpu().numpy(), bs.cpu().numpy()
bs = np.full((N, M), -1e9, "float32"); bi = np.zeros((N, M), np.int64); step = 500000
for i in range(0, g.shape[0], step):
s = Q @ np.asarray(g[i:i+step], dtype="float32").T
ci = np.argpartition(-s, min(M, s.shape[1]-1), 1)[:, :M]; cs = np.take_along_axis(s, ci, 1); ci = ci + i
al = np.concatenate([bs, cs], 1); ai = np.concatenate([bi, ci], 1); o = np.argsort(-al, 1)[:, :M]
bs = np.take_along_axis(al, o, 1); bi = np.take_along_axis(ai, o, 1)
return bi, bs
def zc(a): return (a - a.mean()) / (a.std() + 1e-6)
def cell_key(ll, C): return (np.round(ll[..., 0]/C).astype(np.int64) << 20) ^ (np.round(ll[..., 1]/C).astype(np.int64) & 0xFFFFF)
def predict_one(R, n): # fuse the 3 encoders + density-weighted consensus for image n
sc = {}
for t in ENC_D:
idx, sim = R[t]; z = zc(sim[n])
for j, zz in zip(idx[n], z): sc[int(j)] = sc.get(int(j), 0.0) + float(zz)
ci = np.array(sorted(sc, key=lambda k: -sc[k])[:K], np.int64); csim = np.array([sc[i] for i in ci], "float32")
cc = GC[ci]
DW = np.mean([[-np.log(DENS[C].get(int(k), 1) + 1.0) for k in cell_key(cc, C)] for C in (0.25,0.5,1.0,2.0)], 0)
la, lo = np.radians(cc[:, 0]), np.radians(cc[:, 1])
D = 2*6371.0*np.arcsin(np.sqrt(np.clip(np.sin((la[:,None]-la[None,:])/2)**2 +
np.cos(la)[:,None]*np.cos(la)[None,:]*np.sin((lo[:,None]-lo[None,:])/2)**2, 0, 1)))
vote = np.maximum(csim, 0) * np.exp(Gv * zc(DW))
KV1 = (vote[None,:] * np.exp(-D/BW1)).sum(1); KV2 = (vote[None,:] * np.exp(-D/BW2)).sum(1)
score = zc(csim) + 0.3*zc(DW) + LAM*zc(KV1) + MU*zc(KV2)
return cc[int(score.argmax())]
if __name__ == "__main__":
paths = [p for a in sys.argv[1:] for p in glob.glob(a)] or sys.exit("usage: python run_chipoint.py IMAGE ...")
print(f"device={DEV} loading heads + coords from {W} ...", flush=True)
HEADS = {t: load_head(t, D) for t, D in ENC_D.items()}
GC = np.load(W / "gallery_all_C_so400m256.npy").astype("float32") # GPS per gallery entry
DENS = {C: dict(zip(*[a.tolist() for a in np.unique(cell_key(GC, C), return_counts=True)])) for C in (0.25,0.5,1.0,2.0)}
imgs = [Image.open(p).convert("RGB") for p in paths]
print(f"encoding {len(imgs)} image(s) ...", flush=True)
E = encode_all(imgs) # encoders freed after this
print("retrieving (each gallery read once for the whole batch) ...", flush=True)
R = {t: retrieve_all(t, project(t, E[t])) for t in ENC_D} # one 5 GB pass per encoder
for n, p in enumerate(paths):
lat, lon = predict_one(R, n)
print(f"{p}\t{lat:.5f}, {lon:.5f}\thttps://maps.google.com/?q={lat:.5f},{lon:.5f}", flush=True)