PawTrace / backend /scripts /train_reid.py
Elliott Duke
Cleanup of the system and documentation
d7d2eaf
Raw
History Blame Contribute Delete
31.1 kB
"""Fine-tune a ResNet backbone for dog re-ID with triplet loss. Two interchangeable starting points
(pick with --init); no DINOv2 anywhere:
--init resnet50 torchvision ResNet-50, ImageNet-1k pretrained (~25M params). The standard
person-re-ID backbone: generic features, no breed-invariance baggage, lighter/faster
at inference -- the baseline to beat.
--init breed your jhoppanne/Dogs-Breed-Image-Classification-V1 (ResNet-101, ~44M) with the
classifier head stripped. Dog-domain head start (matters at your scale), and its
training data (Stanford Dogs) is NOT your re-ID set, so validation stays clean --
BUT it was trained toward breed-invariance (suppressing individual differences), so
we unfreeze deep stages to let triplet loss undo that collapse. NOTE: the shipped
checkpoint (best.pt) was trained with ``--unfreeze 1``, not the default below --
diffing it against the pre-trained weights shows only the last stage changed, with
the earlier stages and the stem bit-identical.
Both emit a 2048-d L2-normalized embedding (global-avg-pooled final conv features).
Data (folder-per-identity), produced by scripts/ingest_reid_data.py:
FACE_ROOT/<dog_id>/*.jpg ~1400 face dogs (yours)
body_sources.json ~2800 body dogs (YT-BB-Dog + MPDD + your own)
sibetan_eval_manifest.json leakage-free cross-camera eval
Online random-crop augmentation (random tighter crops labeled as the same identity) trains general
crop/scale invariance -- a cheap mitigation for the face/body framing gap; the Sibetan cross-camera
metric is what tells you whether any of this actually generalizes.
Verify the HF ResNet stage attribute path (backbone.encoder.stages) against your transformers version,
and smoke-test on a tiny subset first.
"""
from __future__ import annotations
import argparse
import random
import time
from dataclasses import dataclass
from pathlib import Path
import torch
import torch.nn as nn
import torch.nn.functional as F
import torchvision.transforms as T
from PIL import Image
# ---- paths -----------------------------------------------------------------------------------
# REID_ROOT defaults to <project>/reid_data locally; on Colab set env REID_ROOT to the unzipped
# bundle. split.json + the sibetan manifest may hold absolute (local) OR bundle-relative paths --
# _resolve() handles both, so the same script runs in both places.
import json
import os
REID_ROOT = Path(os.environ.get("REID_ROOT", Path(__file__).resolve().parents[2] / "reid_data"))
SPLIT_JSON = REID_ROOT / "split.json" # 80/10/10 train/eval/test (build_split.py)
SIBETAN_MANIFEST = REID_ROOT / "sibetan_eval_manifest.json" # separate cross-camera generalization set
# CKPT_DIR: set env CKPT_DIR to a Google Drive folder so checkpoints save straight to Drive DURING
# training -- survives a Colab disconnect (a save-at-the-end approach would lose everything).
CKPT_DIR = Path(os.environ.get("CKPT_DIR", REID_ROOT / "checkpoints"))
def _resolve(p: str) -> Path:
q = Path(p)
return q if q.is_absolute() else (REID_ROOT / q)
# ---- hyperparameters --------------------------------------------------------------------------
INIT = "breed" # "breed" | "resnet50" (overridden by --init)
BREED_CKPT = "jhoppanne/Dogs-Breed-Image-Classification-V1"
UNFREEZE_LAST_N_STAGES = 2 # of 4 ResNet stages. Override with --unfreeze; best.pt used 1.
TRIPLET_MARGIN = 0.45
VAR_LOSS_WEIGHT = 0.5
LR = 1e-4
EPOCHS = 15
STEPS_PER_EPOCH = 200
IDENTITIES_PER_BATCH = 16 # "P" in PK sampling
PHOTOS_PER_IDENTITY = 4 # "K" in PK sampling
DOMAIN_RATIO = 0.5 # fraction of each PK batch drawn from body vs face identities
CROP_AUG_PROB = 0.5 # chance a sampled photo is replaced by a random tighter crop
CROP_AUG_SCALE_RANGE = (0.4, 0.85) # cropped region covers this fraction of width/height
DEVICE = "cuda" if torch.cuda.is_available() else "cpu"
IMAGENET_MEAN, IMAGENET_STD = [0.485, 0.456, 0.406], [0.229, 0.224, 0.225]
_preprocess = T.Compose([T.Resize((224, 224)),
T.ToTensor(),
T.Normalize(IMAGENET_MEAN, IMAGENET_STD)])
# ---- model ---------------------------------------------------------------------------------
class ReIDModel(nn.Module):
"""ResNet backbone -> 2048-d L2-normalized embedding, from one of two starting checkpoints."""
def __init__(self, init: str = INIT):
super().__init__()
self.init = init
if init == "resnet50":
import torchvision
m = torchvision.models.resnet50(weights=torchvision.models.ResNet50_Weights.IMAGENET1K_V2)
m.fc = nn.Identity() # forward now returns the 2048-d pooled features
self.backbone = m
elif init == "breed":
from transformers import AutoModel
# AutoModel loads the ResNet-101 BASE (fine-tuned weights); the classifier head is dropped.
self.backbone = AutoModel.from_pretrained(BREED_CKPT)
else:
raise ValueError(f"unknown init: {init!r}")
def forward(self, images: list[Image.Image]) -> torch.Tensor:
dev = next(self.parameters()).device
x = torch.stack([_preprocess(im.convert("RGB")) for im in images]).to(dev)
if self.init == "resnet50":
feats = self.backbone(x) # (N, 2048)
else:
feats = self.backbone(x).pooler_output.flatten(1) # HF ResNet: (N,2048,1,1)->(N,2048)
return F.normalize(feats, dim=1)
def freeze_backbone(model: ReIDModel, n_stages: int = UNFREEZE_LAST_N_STAGES) -> None:
"""Freeze early stages, unfreeze the last ``n_stages`` (+ their norms) so triplet loss reshapes
only the high-level semantics -- keeps generic edge/texture features, adapts identity."""
for p in model.backbone.parameters():
p.requires_grad = False
if model.init == "resnet50":
stages = [model.backbone.layer1, model.backbone.layer2,
model.backbone.layer3, model.backbone.layer4]
else:
# NOTE: verify against your transformers version -- HF ResNetModel: .encoder.stages (ModuleList).
stages = list(model.backbone.encoder.stages)
for stage in stages[-n_stages:]:
for p in stage.parameters():
p.requires_grad = True
# ---- data -----------------------------------------------------------------------------------
def load_split() -> dict[str, dict[str, list[Path]]]:
"""Load the fixed 80/10/10 split (build_split.py). Keys are 'face:<id>' / 'body:<...>'; paths are
resolved (absolute local paths, or bundle-relative on Colab). Returns {train,val,test}: {key:[paths]}.
Accepts the legacy 'eval' split name and exposes it as 'val' (so old split.json still loads)."""
raw = json.loads(SPLIT_JSON.read_text())
if "eval" in raw and "val" not in raw:
raw["val"] = raw.pop("eval") # legacy name -> validation
return {s: {k: [_resolve(p) for p in paths] for k, paths in d.items()} for s, d in raw.items()}
def by_domain(ids: dict[str, list[Path]], domain: str) -> dict[str, list[Path]]:
return {k: v for k, v in ids.items() if k.startswith(f"{domain}:")}
def load_sibetan_split(holdout_n: int, seed: int = 0) -> tuple[set[str], dict[str, list[Path]]]:
"""Split Sibetan into a held-out cross-camera EVAL set and a TRAIN set. Only multi-camera dogs make
a useful cross-camera eval, so we hold out `holdout_n` of them (seeded, deterministic) and NEVER
train on those. Everything else -- the remaining multi-camera dogs + all single-camera dogs -- is
returned as body-domain training identities (Sibetan is the ONLY source of same-dog/different-scene
positive pairs, exactly the cross-camera invariance we otherwise can't teach).
Returns (holdout_identities, train_ids{ 'body:sibetan__<id>': [paths] })."""
if not SIBETAN_MANIFEST.exists():
return set(), {}
entries = json.loads(SIBETAN_MANIFEST.read_text())
cams: dict[str, set[str]] = {}
imgs: dict[str, list[Path]] = {}
for e in entries:
cams.setdefault(e["identity"], set()).add(e["camera"])
imgs.setdefault(e["identity"], []).append(_resolve(e["path"]))
multi = sorted(i for i, c in cams.items() if len(c) > 1) # deterministic order for the seed
holdout = set(random.Random(seed).sample(multi, min(holdout_n, len(multi)))) if holdout_n else set()
train_ids = {f"body:sibetan__{i}": ps for i, ps in imgs.items() if i not in holdout}
return holdout, train_ids
def random_view(img: Image.Image) -> Image.Image:
"""With CROP_AUG_PROB, return a random tighter crop of the same photo (still the same identity);
otherwise the original. Applied to both domains -- for body photos this occasionally lands on
something close to a face crop by chance; always trains general scale/crop invariance."""
if random.random() > CROP_AUG_PROB:
return img
w, h = img.size
scale = random.uniform(*CROP_AUG_SCALE_RANGE)
cw, ch = int(w * scale), int(h * scale)
x0 = random.randint(0, max(w - cw, 0))
y0 = random.randint(0, max(h - ch, 0))
return img.crop((x0, y0, x0 + cw, y0 + ch))
@dataclass
class PKBatch:
images: list[Image.Image]
identity_ids: list[str] # global-unique across domains (prefixed by domain)
class IdentityPool:
"""An identity->photos map (keys already globally-unique), sampled P identities x K photos."""
def __init__(self, ids: dict[str, list[Path]]):
self.ids = dict(ids)
def sample_batch(self, n_identities: int, k_photos: int) -> PKBatch:
keys = random.sample(list(self.ids), min(n_identities, len(self.ids)))
images, identity_ids = [], []
for key in keys:
photos = self.ids[key]
chosen = random.sample(photos, min(k_photos, len(photos)))
while len(chosen) < k_photos: # identity has fewer than K photos: repeat with replacement
chosen.append(random.choice(photos))
for p in chosen:
images.append(random_view(Image.open(p).convert("RGB")))
identity_ids.append(key)
return PKBatch(images, identity_ids)
def mixed_pk_batch(face_pool: IdentityPool, body_pool: IdentityPool, domain_ratio: float) -> PKBatch:
n_body = round(IDENTITIES_PER_BATCH * domain_ratio)
n_face = IDENTITIES_PER_BATCH - n_body
b = body_pool.sample_batch(n_body, PHOTOS_PER_IDENTITY)
f = face_pool.sample_batch(n_face, PHOTOS_PER_IDENTITY)
return PKBatch(b.images + f.images, b.identity_ids + f.identity_ids)
# ---- losses (triplet + intra-identity variance, same combination as the AvitoTech recipe) ----
def batch_hard_triplet_loss(embeddings: torch.Tensor, identity_ids: list[str],
margin: float = TRIPLET_MARGIN) -> torch.Tensor:
dist = torch.cdist(embeddings, embeddings, p=2)
index = {key: i for i, key in enumerate(dict.fromkeys(identity_ids))}
ids = torch.tensor([index[i] for i in identity_ids], device=embeddings.device)
same = ids.unsqueeze(0) == ids.unsqueeze(1)
diff = ~same
same.fill_diagonal_(False)
hardest_pos = (dist * same).max(dim=1).values
hardest_neg = dist.masked_fill(~diff, float("inf")).min(dim=1).values
return F.relu(hardest_pos - hardest_neg + margin).mean()
def intra_identity_variance_loss(embeddings: torch.Tensor, identity_ids: list[str]) -> torch.Tensor:
total = torch.tensor(0.0, device=embeddings.device)
seen = set(identity_ids)
for key in seen:
idx = [i for i, k in enumerate(identity_ids) if k == key]
if len(idx) < 2:
continue
group = embeddings[idx]
total = total + group.var(dim=0, unbiased=False).mean()
return total / max(len(seen), 1)
def reid_loss(embeddings: torch.Tensor, identity_ids: list[str]) -> torch.Tensor:
return (batch_hard_triplet_loss(embeddings, identity_ids)
+ VAR_LOSS_WEIGHT * intra_identity_variance_loss(embeddings, identity_ids))
# ---- in-domain evaluation: per-modality retrieval (R@k + mAP) on a held-out set ---------------
def retrieval_metrics(E_sub: torch.Tensor, lab_sub: torch.Tensor) -> dict[str, float]:
"""Retrieval quality for a single modality. Every image is a query; the gallery is all OTHER images
in the group (self excluded). R@k = fraction of queries with a same-dog match in the top k. mAP =
mean average precision -- for each query it averages the precision at every rank where a true match
sits, so it rewards pushing ALL of a dog's photos up, not just the first. mAP is far stabler than
R@1 (it moves on any ranking gain), which is why it's the selection/early-stop metric. Only queries
that have >=1 same-dog match are scored."""
sim = E_sub @ E_sub.t()
sim.fill_diagonal_(-2.0) # a query can't retrieve itself
same = lab_sub.unsqueeze(0) == lab_sub.unsqueeze(1)
same.fill_diagonal_(False) # ...nor count itself as a match
has_pos = same.any(dim=1) # only score queries with a match
n = int(has_pos.sum())
if n == 0:
return {"r1": 0.0, "r5": 0.0, "r10": 0.0, "map": 0.0, "n": 0}
order = sim.argsort(dim=1, descending=True) # rank the gallery per query
rel = torch.gather(same, 1, order).float() # 1.0 where the ranked item is a true match
rk = lambda k: round((rel[:, :k].sum(dim=1) > 0)[has_pos].float().mean().item(), 3)
prec = rel.cumsum(dim=1) / torch.arange(1, rel.shape[1] + 1).float() # precision@each rank
ap = (prec * rel).sum(dim=1) / same.sum(dim=1).clamp(min=1) # average precision per query
return {"r1": rk(1), "r5": rk(5), "r10": rk(10),
"map": round(ap[has_pos].mean().item(), 3), "n": n}
@torch.no_grad()
def evaluate_indomain(model: ReIDModel, ids: dict[str, list[Path]], batch: int = 32,
with_loss: bool = False, loss_batches: int = 8) -> dict[str, float]:
"""Held-out (val or test) retrieval, reported separately for faces and bodies: face queries
retrieve the face gallery, body queries the body gallery. Returns R@1/5/10 + mAP per modality.
with_loss=True also returns 'val_loss' -- the SAME reid_loss used in training, averaged over a
FIXED (seeded) set of PK samples drawn from the already-computed embeddings, so it's deterministic
and comparable epoch-to-epoch with no extra forward passes. Rising val_loss while train loss keeps
falling = overfitting."""
model.eval()
paths, labels = [], []
for k, ps in ids.items():
for p in ps:
paths.append(p); labels.append(k)
embs = []
for i in range(0, len(paths), batch):
imgs = [Image.open(p).convert("RGB") for p in paths[i:i + batch]]
embs.append(model(imgs))
E = torch.cat(embs).cpu() # eval math on CPU -> GPU-safe (label tensors are CPU)
order = {k: i for i, k in enumerate(dict.fromkeys(labels))}
lab = torch.tensor([order[l] for l in labels])
is_face = torch.tensor([l.startswith("face:") for l in labels])
def domain(domain_face: bool) -> dict[str, float]:
sel = (is_face == domain_face).nonzero(as_tuple=True)[0]
if len(sel) < 2:
return {"r1": float("nan"), "r5": float("nan"), "r10": float("nan"), "map": float("nan"), "n": 0}
return retrieval_metrics(E[sel], lab[sel])
val_loss = float("nan")
if with_loss:
by_id: dict[int, list[int]] = {}
for i, l in enumerate(lab.tolist()):
by_id.setdefault(l, []).append(i)
usable = sorted(k for k, v in by_id.items() if len(v) >= 2) # need >=2 photos for a positive
rng = random.Random(1234) # fixed -> comparable across epochs
losses = []
for _ in range(loss_batches):
keys = rng.sample(usable, min(IDENTITIES_PER_BATCH, len(usable)))
rows, idl = [], []
for k in keys:
pick = rng.sample(by_id[k], min(PHOTOS_PER_IDENTITY, len(by_id[k])))
rows += pick; idl += [str(k)] * len(pick)
losses.append(reid_loss(E[rows], idl).item())
val_loss = round(sum(losses) / max(len(losses), 1), 4)
model.train()
f, b = domain(True), domain(False)
return {"face_rank1": f["r1"], "face_r5": f["r5"], "face_r10": f["r10"], "face_map": f["map"], "face_n": f["n"],
"body_rank1": b["r1"], "body_r5": b["r5"], "body_r10": b["r10"], "body_map": b["map"], "body_n": b["n"],
"val_loss": val_loss}
@torch.no_grad()
def evaluate_sibetan(model: ReIDModel, manifest_path: Path = SIBETAN_MANIFEST,
batch: int = 32, query_ids: set[str] | None = None) -> dict[str, float]:
"""Leakage-free cross-camera re-ID metric on Sibetan (cross-camera, multi-day camera traps).
For each image, retrieve its nearest neighbor among images from a DIFFERENT track (stricter: a
different camera) and check identity. Excluding same-track candidates removes same-appearance
background cheating, so this measures whether the model learned THE DOG, not the scene.
query_ids (optional): only score queries whose identity is in this set (the held-out cross-camera
dogs, which are never trained on). The gallery still spans ALL Sibetan images, so trained dogs act
as same-domain distractors -- keeps the metric hard AND leakage-free (held-out dogs' images are
never used as training queries/positives). query_ids=None scores every dog (old behavior)."""
if not manifest_path.exists():
return {}
entries = json.loads(manifest_path.read_text())
model.eval()
embs = []
for i in range(0, len(entries), batch):
imgs = [Image.open(_resolve(e["path"])).convert("RGB") for e in entries[i:i + batch]]
embs.append(model(imgs))
E = torch.cat(embs).cpu() # eval math on CPU -> GPU-safe (label tensors are CPU)
sim = E @ E.t()
def code(field: str) -> torch.Tensor:
order = {v: i for i, v in enumerate(dict.fromkeys(e[field] for e in entries))}
return torch.tensor([order[e[field]] for e in entries])
ident, track, camera = code("identity"), code("track"), code("camera")
same_ident = ident.unsqueeze(0) == ident.unsqueeze(1)
# rows we're allowed to SCORE as queries (held-out dogs only, if given); gallery is always all rows.
is_query = (torch.tensor([e["identity"] in query_ids for e in entries])
if query_ids is not None else torch.ones(len(entries), dtype=torch.bool))
def recall(group: torch.Tensor) -> tuple[float, float, float, float, int]:
# gallery excludes the query's own group (same track / same camera) -> no same-scene cheating.
same_group = group.unsqueeze(0) == group.unsqueeze(1)
valid = ~same_group # also excludes self
same = same_ident & valid
has_pos = same.any(dim=1) & is_query # score only held-out queries with a match
n = int(has_pos.sum())
if n == 0:
return 0.0, 0.0, 0.0, 0.0, 0
sim_v = sim.masked_fill(~valid, -2.0) # invalid gallery items ranked last
order = sim_v.argsort(dim=1, descending=True)
rel = torch.gather(same, 1, order).float() # 1 where the ranked item is a true match
rk = lambda k: round((rel[:, :k].sum(dim=1) > 0)[has_pos].float().mean().item(), 3)
prec = rel.cumsum(dim=1) / torch.arange(1, rel.shape[1] + 1).float()
ap = (prec * rel).sum(dim=1) / same.sum(dim=1).clamp(min=1)
return rk(1), rk(5), rk(10), round(ap[has_pos].mean().item(), 3), n
model.train()
xt1, xt5, xt10, xtm, nt = recall(track)
xc1, xc5, xc10, xcm, nc = recall(camera)
return {"xtrack_r1": xt1, "xtrack_r5": xt5, "xtrack_r10": xt10, "xtrack_map": xtm, "xtrack_n": nt,
"xcam_r1": xc1, "xcam_r5": xc5, "xcam_r10": xc10, "xcam_map": xcm, "xcam_n": nc}
# ---- training loop -------------------------------------------------------------------------
def _fmt_secs(s: float) -> str:
m, sec = divmod(int(s), 60); h, m = divmod(m, 60)
return f"{h}h{m:02d}m{sec:02d}s" if h else f"{m}m{sec:02d}s"
def _bucket_table(title: str, ev: dict, sib: dict) -> str:
"""Aligned face/body/track/cam table: mAP + R@1/5/10 + n. (track/cam only if Sibetan present.)"""
rows = [("face", ev["face_map"], ev["face_rank1"], ev["face_r5"], ev["face_r10"], ev["face_n"]),
("body", ev["body_map"], ev["body_rank1"], ev["body_r5"], ev["body_r10"], ev["body_n"])]
if sib:
rows += [("track", sib["xtrack_map"], sib["xtrack_r1"], sib["xtrack_r5"], sib["xtrack_r10"], sib["xtrack_n"]),
("cam", sib["xcam_map"], sib["xcam_r1"], sib["xcam_r5"], sib["xcam_r10"], sib["xcam_n"])]
fnum = lambda v: f"{v:7.3f}" if (isinstance(v, float) and v == v) else f"{'n/a':>7}"
out = [f" {title}", f" {'bucket':<7}{'mAP':>7}{'R@1':>7}{'R@5':>7}{'R@10':>7}{'n':>8}"]
out += [f" {b:<7}{fnum(mp)}{fnum(r1)}{fnum(r5)}{fnum(r10)}{n:>8}" for b, mp, r1, r5, r10, n in rows]
return "\n".join(out)
def train(model: ReIDModel, face_pool: IdentityPool, body_pool: IdentityPool, val_set,
epochs: int, steps: int, eval_every: int, lr: float = LR,
patience: int = 0, min_delta: float = 0.0) -> None:
# Selection/stopping metric = mean(face mAP, body mAP) on the held-out VALIDATION set. mAP is the
# standard re-ID metric and far stabler than R@1, so early-stop decisions are reliable. Improvement
# is: sel > best + min_delta. Only the single best.pt is kept (overwritten on each improvement).
# `patience` > 0 early-stops after that many CONSECUTIVE evals with no improvement.
# Sibetan (cam/track) is NOT evaluated per-epoch -- it's low-res out-of-domain data with no
# influence on the model, so run() records it only at the baseline (start) and the final TEST.
# In practice its numbers turned out not to be usable at all: the crops have a median size of
# roughly 81x75 px against a 224x224 model input, so ~99% of them are upscaled about 9x in area
# before the model sees them. Whatever those scores measure, it is not cross-background
# re-identification, so they are reported for completeness and not treated as a result.
opt = torch.optim.Adam((p for p in model.parameters() if p.requires_grad), lr=lr)
best_sel, best_epoch, stale, n_improved = -1.0, -1, 0, 0
_mean = lambda *xs: sum(x for x in xs if x == x) / max(sum(x == x for x in xs), 1) # nan-safe mean
t_start = time.time()
for epoch in range(epochs):
t_ep = time.time()
model.train()
running = 0.0
for _ in range(steps):
batch = mixed_pk_batch(face_pool, body_pool, DOMAIN_RATIO)
emb = model(batch.images)
loss = reid_loss(emb, batch.identity_ids)
opt.zero_grad(); loss.backward(); opt.step()
running += loss.item()
train_loss = running / steps # epoch-mean training loss
# Eval is the big cost, so only every `eval_every` epochs (and always the last).
if (epoch + 1) % eval_every == 0 or epoch == epochs - 1:
ev = evaluate_indomain(model, val_set, with_loss=True) # face/body only; no Sibetan per-epoch
dt = (time.time() - t_ep) / 60.0 # whole-epoch time (train + eval), minutes
sel = _mean(ev["face_map"], ev["body_map"]) # selection = mean in-domain mAP
line = (f"epoch {epoch:03d} | {dt:.2f} min | "
f"loss tr={train_loss:.3f} val={ev['val_loss']:.3f} | sel(mAP)={sel:.3f} | "
f"body R@1={ev['body_rank1']:.3f} R@5={ev['body_r5']:.3f} | "
f"face R@1={ev['face_rank1']:.3f} R@5={ev['face_r5']:.3f}")
if sel > best_sel + min_delta: # improvement -> keep only the best.pt
best_sel, best_epoch, stale, n_improved = sel, epoch, 0, n_improved + 1
torch.save(model.state_dict(), CKPT_DIR / "best.pt")
print(line + " *SAVED best")
print(_bucket_table(f"val @ new best (epoch {epoch})", ev, None))
else:
stale += 1
print(line + (f" (no improve {stale}/{patience})" if patience else ""))
if patience and stale >= patience:
print(f"\nEARLY STOP: {patience} evals with no improvement "
f"(best sel(mAP)={best_sel:.3f} @ epoch {best_epoch}).")
break
else:
print(f"epoch {epoch:03d} | {(time.time()-t_ep)/60.0:.2f} min | "
f"loss tr={train_loss:.3f} (eval skipped)")
print(f"\nBest sel(mAP)={best_sel:.3f} at epoch {best_epoch} ({n_improved} improvements, only best.pt kept). "
f"Total train time {_fmt_secs(time.time()-t_start)}.")
def run(init: str = INIT, epochs: int = EPOCHS, steps: int = STEPS_PER_EPOCH,
eval_every: int = 1, lr: float = LR, unfreeze: int = UNFREEZE_LAST_N_STAGES,
patience: int = 0, min_delta: float = 0.0,
sibetan_holdout: int = 3, sibetan_train: bool = False) -> None:
CKPT_DIR.mkdir(parents=True, exist_ok=True)
split = load_split()
train_set, val_set, test_set = split["train"], split["val"], split["test"]
nf = lambda s: sum(k.startswith("face:") for k in s)
# Guard: test identities never overlap train/val, so test can't leak into any monitoring decision.
assert set(test_set).isdisjoint(train_set) and set(test_set).isdisjoint(val_set), "test set leaked!"
# Sibetan is its OWN held-out cross-camera set -- never trained on, reported (not gated). Selection
# is gated on the in-domain sel score instead, so xcam is one marker among several, not the only one.
# (Opt-in --sibetan-train folds most of it into training + holds out a few dogs, if you ever want it.)
if sibetan_train:
holdout, sib_train = load_sibetan_split(sibetan_holdout)
sib_query_ids = holdout # score only held-out dogs; gallery = all Sibetan
else:
holdout, sib_train, sib_query_ids = set(), {}, None # held out entirely: eval on all dogs
model = ReIDModel(init=init).to(DEVICE)
freeze_backbone(model, unfreeze)
trainable = sum(p.numel() for p in model.parameters() if p.requires_grad) / 1e6
total = sum(p.numel() for p in model.parameters()) / 1e6
print(f"init={init} device={DEVICE} lr={lr} unfreeze={unfreeze}stages "
f"({trainable:.1f}M/{total:.1f}M trainable) | batch={IDENTITIES_PER_BATCH}ids x "
f"{PHOTOS_PER_IDENTITY} | epochs={epochs} steps={steps}")
face_pool = IdentityPool(by_domain(train_set, "face"))
body_ids = by_domain(train_set, "body")
body_ids.update(sib_train) # add Sibetan (cross-camera) dogs to the body pool
body_pool = IdentityPool(body_ids)
print(f"train {len(train_set)+len(sib_train)} "
f"({nf(train_set)}f/{len(train_set)-nf(train_set)}b + {len(sib_train)} sibetan) "
f"| val {len(val_set)} | test {len(test_set)} | sibetan held-out (xcam): "
f"{sorted(holdout) if holdout else 'none (eval on all)'}")
# Full metrics BEFORE any training (untrained backbone), INCLUDING Sibetan cam/track -- recorded
# here at the start and again at the final TEST, but never during training (see train()).
t_b = time.time()
bev = evaluate_indomain(model, val_set, with_loss=True)
bsib = evaluate_sibetan(model, query_ids=sib_query_ids)
bsel = (bev["face_map"] + bev["body_map"]) / 2
print(f"\n=== BEFORE TRAINING (baseline, untrained) | eval took {(time.time()-t_b)/60.0:.2f} min ===")
print(f"baseline | val loss={bev['val_loss']:.3f} | sel(mAP)={bsel:.3f}")
print(_bucket_table("baseline (val + sibetan)", bev, bsib))
print("=== training starts (Sibetan cam/track NOT evaluated again until the final TEST) ===")
train(model, face_pool, body_pool, val_set, epochs, steps, eval_every, lr, patience, min_delta)
# Final: the untouched TEST set is loaded ONLY here, after training + selection are fully done, so it
# never influenced any monitoring/save/stop decision (train() only ever sees the val set).
model.load_state_dict(torch.load(CKPT_DIR / "best.pt"))
tev = evaluate_indomain(model, test_set)
tsib = evaluate_sibetan(model, query_ids=sib_query_ids)
print(_bucket_table("TEST (best ckpt, held out until now)", tev, tsib))
if __name__ == "__main__":
ap = argparse.ArgumentParser()
ap.add_argument("--init", choices=["breed", "resnet50"], default=INIT,
help="backbone starting checkpoint (default: breed)")
ap.add_argument("--epochs", type=int, default=1000, help="epoch CEILING; patience decides the real stop")
ap.add_argument("--steps", type=int, default=STEPS_PER_EPOCH, help="training steps per epoch")
ap.add_argument("--eval-every", type=int, default=1, help="run eval only every N epochs (CPU saver)")
ap.add_argument("--lr", type=float, default=LR, help="learning rate (lower = finer / less overfit)")
ap.add_argument("--unfreeze", type=int, default=UNFREEZE_LAST_N_STAGES,
help="ResNet stages to unfreeze, 1-4 (fewer = finer / less overfit)")
ap.add_argument("--batch-ids", type=int, default=IDENTITIES_PER_BATCH, help="P: identities per batch")
ap.add_argument("--batch-photos", type=int, default=PHOTOS_PER_IDENTITY, help="K: photos per identity")
ap.add_argument("--patience", type=int, default=30,
help="early-stop after this many evals with no improvement (0 = off, run all epochs)")
ap.add_argument("--min-delta", type=float, default=0.002,
help="minimum sel(mAP) gain to count as an improvement (ignores noise wiggle)")
ap.add_argument("--sibetan-train", action="store_true",
help="opt in: fold most of Sibetan into training, hold out --sibetan-holdout dogs "
"for xcam (default OFF: Sibetan is fully held out, xcam evaluated on all dogs)")
ap.add_argument("--sibetan-holdout", type=int, default=3,
help="with --sibetan-train: # of multi-camera dogs held out for the xcam check")
a = ap.parse_args()
IDENTITIES_PER_BATCH = a.batch_ids # batch = P ids x K photos (read at batch-build time)
PHOTOS_PER_IDENTITY = a.batch_photos
run(a.init, a.epochs, a.steps, a.eval_every, a.lr, a.unfreeze, a.patience, a.min_delta,
a.sibetan_holdout, a.sibetan_train)