#!/usr/bin/env python3 """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): # Magnitude a = feats.norm(dim=-1, keepdim=True).clamp(self.l_a, self.u_a) # Margin interpolation m = self.l_m + (self.u_m - self.l_m) * (a - self.l_a) / (self.u_a - self.l_a) # Regularizer g(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() # (B, C) if target is None: return cos * self.s, g # Apply margin only to true class 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 # (B, C), broadcast margin 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()