#!/usr/bin/env python3 """Evaluate the trained VSL T2M-GPT on a held-out split. Reports, for the requested split: * tokenizer ceiling -- MPJPE of VQ-VAE reconstruction (encode->decode of GT). No generation can beat this, so every generation number should be read against it. * teacher-forced -- CE loss and next-token accuracy. * generation -- prefix MPJPE, DTW-MJE (alignment-free, so a length/timing mismatch is not double-counted), and length ratio. DTW-MJE is included because this project's earlier BARTpho gloss->pose baseline was reported in that metric (hand 0.237 / body 0.228 vs a 0.012 ceiling), and the two numbers are only comparable if computed the same way. """ import argparse import json import os import numpy as np import torch from tqdm import tqdm import models.t2m_trans as trans from dataset import dataset_vsl from models.text_encoder_vi import ViTextEncoder from train_t2m_trans_vsl import build_vqvae GROUPS = {"all": (0, 128), "body": (0, 18), "face": (18, 86), "hands": (86, 128)} # replaced at runtime from the data layout NK = 128 # replaced at runtime from the data layout def per_frame_dist(a, b): """a [T1,nk,2], b [T2,nk,2] -> pairwise per-keypoint distances [T1,T2,nk].""" return np.linalg.norm(a[:, None] - b[None, :], axis=-1) def dtw_mje(pred, gt, valid_gt): """Alignment-free mean joint error per keypoint group. pred [T1,128,2], gt [T2,128,2], valid_gt [T2,128]. DTW path is found on the all-keypoint cost, then each group is averaged along that single shared path (so groups stay comparable to each other). """ T1, T2 = len(pred), len(gt) d = per_frame_dist(pred, gt) # [T1,T2,128] w = valid_gt[None, :, :] # [1,T2,128] cost = (d * w).sum(-1) / np.maximum(w.sum(-1), 1e-6) # [T1,T2] # standard DTW with the usual 3 moves D = np.full((T1 + 1, T2 + 1), np.inf) D[0, 0] = 0.0 for i in range(1, T1 + 1): ci = cost[i - 1] for j in range(1, T2 + 1): D[i, j] = ci[j - 1] + min(D[i - 1, j], D[i, j - 1], D[i - 1, j - 1]) # backtrack path, i, j = [], T1, T2 while i > 0 and j > 0: path.append((i - 1, j - 1)) step = int(np.argmin([D[i - 1, j - 1], D[i - 1, j], D[i, j - 1]])) if step == 0: i, j = i - 1, j - 1 elif step == 1: i -= 1 else: j -= 1 pi = np.array([p[0] for p in path]) pj = np.array([p[1] for p in path]) out = {} for name, (a, b) in GROUPS.items(): dd = d[pi, pj, a:b] vv = valid_gt[pj, a:b] out[name] = float((dd * vv).sum() / max(vv.sum(), 1e-6)) return out def prefix_mje(pred, gt, valid_gt): T = min(len(pred), len(gt)) d = np.linalg.norm(pred[:T] - gt[:T], axis=-1) # [T,128] v = valid_gt[:T] out = {} for name, (a, b) in GROUPS.items(): out[name] = float((d[:, a:b] * v[:, a:b]).sum() / max(v[:, a:b].sum(), 1e-6)) return out def main(): ap = argparse.ArgumentParser() ap.add_argument('--data-dir', default='./dataset/VSL') ap.add_argument('--token-dir', default='./dataset/VSL/tokens') ap.add_argument('--resume-pth', required=True, help='stage-1 VQ-VAE checkpoint') ap.add_argument('--resume-trans', required=True, help='stage-2 GPT checkpoint') ap.add_argument('--split', default='test') ap.add_argument('--n', type=int, default=300, help='clips to score (0 = all)') ap.add_argument('--device', default='cuda') ap.add_argument('--sampling', default='categorial', choices=['greedy', 'categorial'], help="categorial matches upstream evaluation_transformer_test") ap.add_argument('--shuffle-text', action='store_true', help='CONTROL: condition each clip on another clip\'s gloss. If this ' 'scores the same as the real pairing, the model is ignoring the text.') ap.add_argument('--text-override', default=None, help='JSON {clip_name: {"pred_gloss": ...}} -- condition on PREDICTED ' 'gloss instead of ground truth, for end-to-end text->pose eval') ap.add_argument('--out-json', default=None) args = ap.parse_args() override = {} if args.text_override: with open(args.text_override, encoding='utf-8') as f: raw = json.load(f) override = {k: (v['pred_gloss'] if isinstance(v, dict) else v) for k, v in raw.items()} print(f'text-override: {len(override)} clips from {args.text_override}') device = torch.device(args.device) net, targs, ck1 = build_vqvae(args.resume_pth, device) tck = torch.load(args.resume_trans, map_location='cpu') gargs = argparse.Namespace(**tck['args']) text_enc = ViTextEncoder(gargs.text_model, device=args.device) trans_encoder = trans.Text2Motion_Transformer( num_vq=gargs.nb_code, embed_dim=gargs.embed_dim_gpt, clip_dim=text_enc.dim, block_size=gargs.max_tokens + 1, num_layers=gargs.num_layers, n_head=gargs.n_head_gpt, drop_out_rate=gargs.drop_out_rate, fc_rate=gargs.ff_rate) trans_encoder.load_state_dict(tck['trans'], strict=True) trans_encoder.eval().to(device) print(f"stage-2 ckpt iter {tck.get('iter')} val_loss {tck.get('val_loss'):.4f} " f"val_acc {tck.get('val_acc'):.2f}") store = dataset_vsl.VSLStore(args.data_dir, args.split) global GROUPS, NK GROUPS = store.layout.metric_groups() NK = store.layout.n_kpts print(f'layout: {store.layout}') rng = np.random.RandomState(0) n = len(store.index) if not args.n else min(args.n, len(store.index)) items = sorted(rng.choice(len(store.index), size=n, replace=False).tolist()) mean, std = store.mean, store.std mean_t = torch.from_numpy(mean).to(device) std_t = torch.from_numpy(std).to(device) agg = {k: {g: [] for g in GROUPS} for k in ('ceiling', 'gen_prefix', 'gen_dtw')} ratios, empties, n_nooverride = [], 0, 0 with torch.no_grad(): for i in tqdm(items, desc=f'eval {args.split}'): c = store.index[i] motion, mask = store.get(i) gt_t = torch.from_numpy(motion).unsqueeze(0).to(device) valid = mask[:, ::2] # [T,128] gt_xy = (motion * std + mean).reshape(-1, NK, 2) # --- tokenizer ceiling: encode then decode the ground truth --- unit = 2 ** targs.down_t T = (len(motion) // unit) * unit codes = net.encode(gt_t[:, :T]) rec = net.decode_batch(codes)[0].cpu().numpy() * std + mean rec = rec.reshape(-1, NK, 2) for g, v in prefix_mje(rec, gt_xy, valid).items(): agg['ceiling'][g].append(v) # --- generation from gloss --- cond_text = c[gargs.text_field] if override: if c['name'] not in override: n_nooverride += 1 continue cond_text = override[c['name']] if args.shuffle_text: # deterministic mismatch: pair clip k with the gloss of another clip other = store.index[(i + len(store.index) // 2) % len(store.index)] cond_text = other[gargs.text_field] feat = text_enc([cond_text]) idx = trans_encoder.sample(feat, if_categorial=(args.sampling == 'categorial')) if idx is None or idx.numel() == 0: empties += 1 continue idx = idx.clamp(max=gargs.nb_code - 1) pred = net.decode_batch(idx)[0].cpu().numpy() * std + mean pred = pred.reshape(-1, NK, 2) ratios.append(len(pred) / len(gt_xy)) for g, v in prefix_mje(pred, gt_xy, valid).items(): agg['gen_prefix'][g].append(v) for g, v in dtw_mje(pred, gt_xy, valid).items(): agg['gen_dtw'][g].append(v) res = {'split': args.split, 'n_clips': n, 'n_empty': empties, 'n_missing_override': n_nooverride, 'text_source': 'predicted' if override else 'ground_truth', 'sampling': args.sampling, 'len_ratio_mean': float(np.mean(ratios)) if ratios else 0.0, 'stage1': {'iter': ck1.get('iter'), 'codes_used': ck1.get('codes_used')}, 'stage2': {'iter': tck.get('iter'), 'val_acc': tck.get('val_acc')}} for k, d in agg.items(): res[k] = {g: (float(np.mean(v)) if v else None) for g, v in d.items()} print(json.dumps(res, indent=2)) print("\n--- summary (frame-normalized units, lower is better) ---") print(f"{'metric':<22}{'all':>9}{'body':>9}{'face':>9}{'hands':>9}") for k in ('ceiling', 'gen_prefix', 'gen_dtw'): row = ''.join(f"{res[k][g]:>9.4f}" if res[k][g] is not None else f"{'-':>9}" for g in ('all', 'body', 'face', 'hands')) print(f"{k:<22}{row}") print(f"len_ratio (pred/gt): {res['len_ratio_mean']:.3f} empty gens: {empties}") if args.out_json: os.makedirs(os.path.dirname(args.out_json) or '.', exist_ok=True) with open(args.out_json, 'w') as f: json.dump(res, f, indent=2) print(f"\nwrote {args.out_json}") if __name__ == '__main__': main()