| |
| """MagFace (Meng CVPR 2021) — magnitude-aware angular margin. |
| |
| Unlike ArcFace (fixed margin m), MagFace scales margin by feature magnitude. |
| High-quality (confident) samples get larger margin → tighter intra-class |
| clustering. Low-quality (tail/ambiguous) samples get smaller margin. |
| |
| Loss: |
| m(||f||) = u_m * (||f|| - l_a) / (u_a - l_a) (interpolation) |
| g(||f||) = (1/l_a - 1/u_m) * ||f|| (regularizer to push magnitudes) |
| """ |
| import os, sys, json, argparse, time, math |
| from pathlib import Path |
| from collections import Counter |
|
|
| import numpy as np |
| import torch |
| import torch.nn as nn |
| import torch.nn.functional as F |
| from torch.utils.data import DataLoader, WeightedRandomSampler |
|
|
| ROOT = Path("/arf/scratch/stakan/hitit-proje") |
| sys.path.insert(0, str(ROOT / "hitit_ocr/src")) |
| from train_classification import build_backbone, get_arch_img_size, HititClsDataset |
|
|
| def log(m): print(f"[{time.strftime('%H:%M:%S')}] {m}", flush=True) |
|
|
|
|
| class MagFaceHead(nn.Module): |
| def __init__(self, feat_dim, n_classes, s=64.0, |
| l_a=10.0, u_a=110.0, l_m=0.45, u_m=0.8, lambda_g=20.0): |
| super().__init__() |
| self.W = nn.Parameter(torch.randn(n_classes, feat_dim)) |
| nn.init.xavier_uniform_(self.W) |
| self.s = s; self.l_a = l_a; self.u_a = u_a |
| self.l_m = l_m; self.u_m = u_m |
| self.lambda_g = lambda_g |
| self.n_classes = n_classes |
|
|
| def forward(self, feats, target=None): |
| |
| a = feats.norm(dim=-1, keepdim=True).clamp(self.l_a, self.u_a) |
| |
| m = self.l_m + (self.u_m - self.l_m) * (a - self.l_a) / (self.u_a - self.l_a) |
| |
| g = (1.0 / self.l_a - 1.0 / self.u_a) * a - (1.0 / self.l_a - 1.0 / self.u_a) * self.l_a |
| g = g.mean() |
|
|
| f = F.normalize(feats, dim=-1) |
| w = F.normalize(self.W, dim=-1) |
| cos = f @ w.t() |
| if target is None: |
| return cos * self.s, g |
|
|
| |
| one_hot = F.one_hot(target, self.n_classes).float() |
| theta = torch.acos(cos.clamp(-1 + 1e-7, 1 - 1e-7)) |
| theta_m = theta + one_hot * m |
| logits = torch.cos(theta_m) * self.s |
| return logits, g |
|
|
|
|
| @torch.no_grad() |
| def extract(model, x, dtype): |
| raw = model.module if hasattr(model, 'module') else model |
| with torch.amp.autocast('cuda', dtype=dtype, enabled=True): |
| if hasattr(raw, 'forward_features'): |
| f = raw.forward_features(x) |
| if f.dim() == 3: f = f[:, 0] |
| elif f.dim() == 4: f = f.mean(dim=(2, 3)) |
| else: f = raw(x) |
| return f.float() |
|
|
|
|
| def main(): |
| ap = argparse.ArgumentParser() |
| ap.add_argument('--ckpt', required=True) |
| ap.add_argument('--manifest', required=True) |
| ap.add_argument('--val-fold', type=int, default=0) |
| ap.add_argument('--min-samples', type=int, default=10) |
| ap.add_argument('--epochs', type=int, default=30) |
| ap.add_argument('--lr', type=float, default=1e-2) |
| ap.add_argument('--s', type=float, default=64.0) |
| ap.add_argument('--lambda-g', type=float, default=20.0) |
| ap.add_argument('--batch-size', type=int, default=128) |
| ap.add_argument('--output', required=True) |
| args = ap.parse_args() |
|
|
| device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') |
| dtype = torch.bfloat16 if torch.cuda.is_available() and torch.cuda.is_bf16_supported() else torch.float32 |
|
|
| arch, path = args.ckpt.split(':', 1) |
| ck = torch.load(path, map_location='cpu', weights_only=False) |
| label_to_idx = ck['label_to_idx'] |
| n_cls = len(label_to_idx) |
| img_size = get_arch_img_size(arch) |
| backbone = build_backbone(arch, n_classes=n_cls, img_size_override=img_size).to(device) |
| sd = ck['model'] |
| sd = {k.replace('module.', '', 1): v for k, v in sd.items()} |
| sd = {k.replace('_orig_mod.', '', 1): v for k, v in sd.items()} |
| sd = {k: v for k, v in sd.items() if k != 'n_averaged'} |
| backbone.load_state_dict(sd, strict=False) |
| for p in backbone.parameters(): p.requires_grad = False |
| backbone.eval() |
|
|
| cfg = {'img_size': img_size} |
| tr_ds = HititClsDataset(args.manifest, cfg, is_train=True, |
| val_fold=args.val_fold, label_to_idx=label_to_idx, |
| min_samples=args.min_samples) |
| va_ds = HititClsDataset(args.manifest, cfg, is_train=False, |
| val_fold=args.val_fold, label_to_idx=label_to_idx, |
| min_samples=args.min_samples) |
| cc = Counter(tr_ds.label_to_idx[r['unified_label']] for r in tr_ds.records) |
| sample_w = np.array([1.0 / max(1, cc[tr_ds.label_to_idx[r['unified_label']]])**0.5 |
| for r in tr_ds.records], dtype=np.float32) |
| tr_dl = DataLoader(tr_ds, batch_size=args.batch_size, |
| sampler=WeightedRandomSampler(sample_w, len(tr_ds), True), |
| num_workers=6, pin_memory=True, drop_last=True) |
| va_dl = DataLoader(va_ds, batch_size=args.batch_size, shuffle=False, |
| num_workers=4, pin_memory=True) |
|
|
| with torch.no_grad(): |
| feat_dim = extract(backbone, torch.zeros(1, 3, img_size, img_size, device=device), dtype).size(-1) |
| head = MagFaceHead(feat_dim, n_cls, s=args.s, lambda_g=args.lambda_g).to(device) |
| opt = torch.optim.AdamW(head.parameters(), lr=args.lr, weight_decay=1e-4) |
| sched = torch.optim.lr_scheduler.CosineAnnealingLR(opt, T_max=args.epochs) |
|
|
| best, best_probs, best_y = 0, None, None |
| for ep in range(args.epochs): |
| head.train() |
| tl, nb = 0, 0 |
| for x, y in tr_dl: |
| x = x.to(device, non_blocking=True); y = y.to(device, non_blocking=True) |
| f = extract(backbone, x, dtype) |
| logits, g = head(f, target=y) |
| loss_ce = F.cross_entropy(logits, y, label_smoothing=0.1) |
| loss = loss_ce + args.lambda_g * g |
| opt.zero_grad(); loss.backward(); opt.step() |
| tl += loss.item(); nb += 1 |
| sched.step() |
| head.eval() |
| cs, tot, p_all, y_all = 0, 0, [], [] |
| with torch.no_grad(): |
| for x, y in va_dl: |
| x = x.to(device, non_blocking=True); y = y.to(device, non_blocking=True) |
| f = extract(backbone, x, dtype) |
| logits, _ = head(f) |
| p = F.softmax(logits.float(), dim=-1) |
| p_all.append(p.cpu()); y_all.append(y.cpu()) |
| cs += (logits.argmax(-1) == y).sum().item(); tot += y.size(0) |
| acc = cs / max(1, tot) |
| if ep % 5 == 0 or ep == args.epochs - 1: |
| log(f"ep {ep}: loss={tl/max(1,nb):.4f} val_top1={acc:.4f}") |
| if acc > best: |
| best = acc |
| best_probs = torch.cat(p_all); best_y = torch.cat(y_all) |
| torch.save({'head_state': head.state_dict(), 'label_to_idx': label_to_idx, |
| 'probs': best_probs, 'targets': best_y, 'top1': best, |
| 'backbone_arch': arch, 's': args.s}, |
| args.output) |
| log(f"=== MagFace BEST: {best:.4f}") |
|
|
| if __name__ == '__main__': |
| main() |
|
|