t2m-gpt-vsl-code / eval_fgd_maej.py
Tri1's picture
T2M-GPT VSL adaptation: Python sources only (82 files, no checkpoints or data)
8e5456b verified
Raw
History Blame Contribute Delete
10.1 kB
#!/usr/bin/env python3
"""FGD and MAEJ (Hwang et al., NSLP-G / BMVC) for every dumped Multi-VSL condition.
Definitions, from the paper (`0.NSLP-G/1102.pdf`, eq. 9 and the surrounding text):
FGD -- Frechet distance between Gaussians fitted to LATENT FEATURES of real vs
generated pose sequences:
||mu_r - mu_g||^2 + Tr(Sig_r + Sig_g - 2 (Sig_r Sig_g)^(1/2))
"FGD measures the diversity of produced sign poses". It needs a feature
extractor; the paper uses a Transformer autoencoder trained on sign poses,
which is what --train-ae builds here (on the GT TRAIN split only, so it
never sees generated poses).
MAEJ -- "independently measures the distance between the produced sign pose and
real sign pose (NO ASPECT OF TEMPORAL DISTANCE)". So it is a per-joint mean
absolute error on coordinates with no DTW: the prediction is resampled to
the reference length and compared frame-to-frame. The paper reports
MAEJ* = MAEJ x 100.
Both are computed on the shared 50-joint layout (8 body + 42 hands), and every
condition is matched to the ground truth BY CLIP NAME, because the two projects
select their 300-clip subsets with different RNGs.
FGD is distributional, so it is computed over whatever clips a condition has; MAEJ is
per-clip and averaged. Neither uses DTW -- that is the point: they probe different
failure modes than the DTW-MJE table.
"""
import argparse
import glob
import json
import os
import numpy as np
import torch
import torch.nn as nn
from scipy import linalg
# ---------------------------------------------------------------- FGD core
def frechet(mu1, s1, mu2, s2, eps=1e-6):
"""Standard Frechet distance (same formula as 0.NSLP-G/modules/fid.py)."""
diff = mu1 - mu2
covmean, _ = linalg.sqrtm(s1.dot(s2), disp=False)
if not np.isfinite(covmean).all():
off = np.eye(s1.shape[0]) * eps
covmean = linalg.sqrtm((s1 + off).dot(s2 + off))
if np.iscomplexobj(covmean):
covmean = covmean.real
return float(diff.dot(diff) + np.trace(s1) + np.trace(s2) - 2 * np.trace(covmean))
def mean_cov(z):
z = np.asarray(z, np.float64)
return z.mean(0), np.cov(z, rowvar=False)
# ------------------------------------------------------- Transformer AE (features)
class PosEnc(nn.Module):
def __init__(self, d, n=1024):
super().__init__()
pe = torch.zeros(n, d)
pos = torch.arange(n).unsqueeze(1).float()
div = torch.exp(torch.arange(0, d, 2).float() * (-np.log(10000.0) / d))
pe[:, 0::2] = torch.sin(pos * div)
pe[:, 1::2] = torch.cos(pos * div)
self.register_buffer('pe', pe.unsqueeze(0))
def forward(self, x):
return x + self.pe[:, :x.shape[1]]
class TFAE(nn.Module):
"""Transformer autoencoder; the clip feature is the time-mean of encoder states."""
def __init__(self, dim_in, d=256, heads=4, layers=3, ff=1024):
super().__init__()
self.inp = nn.Linear(dim_in, d)
self.pe = PosEnc(d)
el = nn.TransformerEncoderLayer(d, heads, ff, batch_first=True, dropout=0.1)
self.enc = nn.TransformerEncoder(el, layers)
dl = nn.TransformerEncoderLayer(d, heads, ff, batch_first=True, dropout=0.1)
self.dec = nn.TransformerEncoder(dl, layers)
self.out = nn.Linear(d, dim_in)
def encode(self, x, mask=None):
h = self.enc(self.pe(self.inp(x)), src_key_padding_mask=mask)
if mask is None:
return h.mean(1)
w = (~mask).float().unsqueeze(-1)
return (h * w).sum(1) / w.sum(1).clamp(min=1)
def forward(self, x, mask=None):
h = self.enc(self.pe(self.inp(x)), src_key_padding_mask=mask)
return self.out(self.dec(h, src_key_padding_mask=mask))
def pad_batch(seqs, device, maxlen=256):
L = min(max(len(s) for s in seqs), maxlen)
D = seqs[0].shape[-1]
x = np.zeros((len(seqs), L, D), np.float32)
m = np.ones((len(seqs), L), bool)
for i, s in enumerate(seqs):
n = min(len(s), L)
x[i, :n] = s[:n]
m[i, :n] = False
return torch.from_numpy(x).to(device), torch.from_numpy(m).to(device)
def flat(p):
return np.asarray(p, np.float32).reshape(len(p), -1)
def train_ae(train_seqs, device, dim_in, epochs=30, bs=64, lr=3e-4):
ae = TFAE(dim_in).to(device)
opt = torch.optim.AdamW(ae.parameters(), lr=lr, weight_decay=1e-4)
idx = np.arange(len(train_seqs))
for ep in range(1, epochs + 1):
np.random.shuffle(idx)
tot = n = 0
ae.train()
for i in range(0, len(idx) - bs + 1, bs):
seqs = [train_seqs[j] for j in idx[i:i + bs]]
x, m = pad_batch(seqs, device)
rec = ae(x, m)
w = (~m).float().unsqueeze(-1)
loss = (((rec - x) ** 2) * w).sum() / w.sum().clamp(min=1) / x.shape[-1]
opt.zero_grad(); loss.backward(); opt.step()
tot += loss.item(); n += 1
if ep % 5 == 0 or ep == 1:
print(f' AE epoch {ep:3d} recon {tot/max(n,1):.5f}')
ae.eval()
return ae
@torch.no_grad()
def features(ae, seqs, device, bs=64):
out = []
for i in range(0, len(seqs), bs):
x, m = pad_batch(seqs[i:i + bs], device)
out.append(ae.encode(x, m).cpu().numpy())
return np.concatenate(out, 0)
# ------------------------------------------------------------------- MAEJ
def resample(p, T):
if len(p) == T:
return p.astype(np.float32)
src, dst = np.linspace(0, 1, len(p)), np.linspace(0, 1, T)
f = p.reshape(len(p), -1)
o = np.stack([np.interp(dst, src, f[:, k]) for k in range(f.shape[1])], 1)
return o.reshape(T, *p.shape[1:]).astype(np.float32)
def maej(pred, gt):
"""Mean absolute joint error, no temporal alignment (pred resampled to len(gt))."""
return float(np.abs(resample(pred, len(gt)) - gt).mean())
# ------------------------------------------------------------------- driver
def load_dump(p):
d = np.load(p, allow_pickle=True)
return {n: q for n, q in zip(d['names'], d['poses'])}
def main():
ap = argparse.ArgumentParser()
ap.add_argument('--dump-dirs', nargs='+',
default=['dumps_fgd_t2m',
'../0.NSLP-G/Word-level/NSLP-G/dumps_fgd'])
ap.add_argument('--gt-npz', default='dumps_fgd_t2m/gt.npz')
ap.add_argument('--data-dir', default='./dataset/MVSL')
ap.add_argument('--epochs', type=int, default=30)
ap.add_argument('--device', default='cuda')
ap.add_argument('--out-json', default='output_vsl/fgd_maej.json')
args = ap.parse_args()
device = torch.device(args.device)
gt = load_dump(args.gt_npz)
print(f'ground truth: {len(gt)} clips')
# --- feature extractor: trained on GT TRAIN poses only ---
from dataset import dataset_vsl
from dump_mvsl_poses import KEEP_50
tr = dataset_vsl.VSLStore(args.data_dir, 'train')
NK = tr.layout.n_kpts
train_seqs = []
for i in range(len(tr.index)):
m, _ = tr.get(i)
train_seqs.append(flat((m * tr.std + tr.mean).reshape(-1, NK, 2)[:, KEEP_50]))
print(f'AE training on {len(train_seqs)} GT train clips, dim {train_seqs[0].shape[-1]}')
ae = train_ae(train_seqs, device, train_seqs[0].shape[-1], epochs=args.epochs)
# --- collect conditions ---
conds = {}
for d in args.dump_dirs:
for p in sorted(glob.glob(os.path.join(d, '*.npz'))):
tag = os.path.splitext(os.path.basename(p))[0]
if tag == 'gt':
continue
src = 'T2M-GPT' if 'dumps_fgd_t2m' in d else 'NSLP-G'
conds[f'{src}:{tag}'] = load_dump(p)
# FGD is biased by sample size, so every condition MUST be scored on the same
# clips: NSLP-G dumped 300 while this project dumps the whole split. Restrict to
# the intersection over all conditions (and the GT), or the numbers are not
# comparable to each other at all.
common_all = set(gt)
for m in conds.values():
common_all &= set(m)
gnames = sorted(common_all)
print(f'\nscoring every condition on the SAME {len(gnames)} clips '
f'(intersection over {len(conds)} conditions + GT)')
for tag, m in sorted(conds.items()):
print(f' {tag:<34} dumped {len(m)}')
if len(gnames) < 20:
raise SystemExit('too few clips shared across conditions')
# real-vs-real reference: split GT in half to expose the FGD noise floor
half = len(gnames) // 2
fa = features(ae, [flat(gt[n]) for n in gnames[:half]], device)
fb = features(ae, [flat(gt[n]) for n in gnames[half:]], device)
ref_fgd = frechet(*mean_cov(fa), *mean_cov(fb))
gt_all = features(ae, [flat(gt[n]) for n in gnames], device)
mu_r, sig_r = mean_cov(gt_all)
rows = []
for tag, m in sorted(conds.items()):
common = gnames # identical clip set for every condition
fz = features(ae, [flat(m[n]) for n in common], device)
fgd = frechet(mu_r, sig_r, *mean_cov(fz))
mj = float(np.mean([maej(np.asarray(m[n], np.float32),
np.asarray(gt[n], np.float32)) for n in common]))
lr = float(np.mean([len(m[n]) / len(gt[n]) for n in common]))
rows.append({'condition': tag, 'n': len(common), 'FGD': fgd,
'MAEJ': mj, 'MAEJ_x100': mj * 100, 'len_ratio': lr})
rows.sort(key=lambda r: r['MAEJ'])
print(f"\n{'condition':<34}{'n':>5}{'FGD':>10}{'MAEJ*':>9}{'len_r':>8}")
print(f"{'real vs real (FGD noise floor)':<34}{half:>5}{ref_fgd:>10.3f}{0.0:>9.3f}{1.0:>8.3f}")
for r in rows:
print(f"{r['condition']:<34}{r['n']:>5}{r['FGD']:>10.3f}"
f"{r['MAEJ_x100']:>9.3f}{r['len_ratio']:>8.3f}")
os.makedirs(os.path.dirname(args.out_json) or '.', exist_ok=True)
with open(args.out_json, 'w') as f:
json.dump({'fgd_noise_floor': ref_fgd, 'n_gt': len(gnames), 'rows': rows},
f, indent=2)
print(f'\nwrote {args.out_json}')
if __name__ == '__main__':
main()