savastakan's picture
Initial upload: code + 5 record checkpoints + fuse
f211247 verified
Raw
History Blame Contribute Delete
7.17 kB
#!/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()