#!/usr/bin/env python3 """T2M-GPT stage 1 on Full_TriVis: VQ-VAE pose tokenizer for DWPose skeletons. Same network and training recipe as the paper's `train_vq.py` (Resnet1D encoder/decoder, EMA + code-reset quantizer, recon + commit + velocity loss). Two things had to change for this data: * motion is 256-dim DWPose xy instead of 263-dim HumanML3D features, and the reconstruction loss is masked by keypoint validity and weighted toward hands; * the paper evaluates with HumanML3D's pretrained FID/R-precision evaluators, which do not exist for VSL, so validation reports masked reconstruction loss and MPJPE (in frame-normalized units) per keypoint group plus codebook usage. """ import json import os import numpy as np import torch import torch.optim as optim from torch.utils.tensorboard import SummaryWriter import models.vqvae as vqvae import options.option_vsl as option_vsl import utils.utils_model as utils_model from dataset import dataset_vsl from utils.losses_vsl import VSLReConsLoss, mpjpe_groups def update_lr_warm_up(optimizer, nb_iter, warm_up_iter, lr): current_lr = lr * (nb_iter + 1) / (warm_up_iter + 1) for param_group in optimizer.param_groups: param_group["lr"] = current_lr return optimizer, current_lr @torch.no_grad() def evaluate(net, loader, Loss, mean_t, std_t, device, groups=None, n_kpts=128): net.eval() tot_recon, tot_vel, nb = 0.0, 0.0, 0 mp = {"all": 0.0, "body": 0.0, "face": 0.0, "hands": 0.0} used = torch.zeros(net.vqvae.num_code, dtype=torch.bool, device=device) for motion, mask in loader: motion = motion.to(device).float() mask = mask.to(device).float() pred, loss_commit, _ = net(motion) tot_recon += Loss(pred, motion, mask).item() tot_vel += Loss.forward_vel(pred, motion, mask).item() # MPJPE in raw frame-normalized coordinates (comparable across runs) pred_xy = pred * std_t + mean_t gt_xy = motion * std_t + mean_t valid = mask.view(mask.shape[0], mask.shape[1], n_kpts, 2)[..., 0] g = mpjpe_groups(pred_xy, gt_xy, valid, groups=groups) for k in mp: mp[k] += g[k] used[net.encode(motion).reshape(-1).unique()] = True nb += 1 net.train() nb = max(nb, 1) return (tot_recon / nb, tot_vel / nb, {k: v / nb for k, v in mp.items()}, int(used.sum().item())) def main(): args = option_vsl.get_vq_args() torch.manual_seed(args.seed) np.random.seed(args.seed) device = torch.device(args.device) args.out_dir = os.path.join(args.out_dir, args.exp_name) os.makedirs(args.out_dir, exist_ok=True) logger = utils_model.get_logger(args.out_dir) writer = SummaryWriter(args.out_dir) logger.info(json.dumps(vars(args), indent=4, sort_keys=True)) ##### ---- Dataloaders ---- ##### train_set = dataset_vsl.VSLVQDataset(args.data_dir, 'train', window_size=args.window_size) train_loader = torch.utils.data.DataLoader( train_set, args.batch_size, shuffle=True, num_workers=args.num_workers, drop_last=True, pin_memory=True, persistent_workers=args.num_workers > 0) train_loader_iter = dataset_vsl.cycle(train_loader) val_set = dataset_vsl.VSLFixedWindowDataset( args.data_dir, 'val', window_size=args.window_size, stride=args.window_size, max_windows=args.val_windows) val_loader = torch.utils.data.DataLoader( val_set, args.batch_size, shuffle=False, num_workers=4, drop_last=False) mean_t = torch.from_numpy(train_set.store.mean).to(device) std_t = torch.from_numpy(train_set.store.std).to(device) layout = train_set.store.layout groups = layout.metric_groups() # the packed data dictates the model's input dim -- never trust the CLI default if args.input_dim != layout.dim: logger.info(f'--input-dim {args.input_dim} overridden by data layout -> {layout.dim}') args.input_dim = layout.dim args.layout = layout.name logger.info(f'layout: {layout}') ##### ---- Network ---- ##### net = vqvae.HumanVQVAE(args, args.nb_code, args.code_dim, args.output_emb_width, args.down_t, args.stride_t, args.width, args.depth, args.dilation_growth_rate, args.vq_act, args.vq_norm, input_dim=args.input_dim) if args.resume_pth: logger.info(f'loading checkpoint from {args.resume_pth}') ckpt = torch.load(args.resume_pth, map_location='cpu') net.load_state_dict(ckpt['net'], strict=True) net.train().to(device) logger.info(f'VQ-VAE params: {sum(p.numel() for p in net.parameters())/1e6:.2f}M, ' f'input_dim={net.vqvae.input_dim}, tokens per {args.window_size} frames = ' f'{args.window_size // 2**args.down_t}') optimizer = optim.AdamW(net.parameters(), lr=args.lr, betas=(0.9, 0.99), weight_decay=args.weight_decay) scheduler = torch.optim.lr_scheduler.MultiStepLR(optimizer, milestones=args.lr_scheduler, gamma=args.gamma) Loss = VSLReConsLoss(args.recons_loss, args.w_body, args.w_face, args.w_hand, layout=layout, w_finger=args.w_finger, w_fingertip=args.w_fingertip).to(device) def step(motion, mask): pred, loss_commit, perplexity = net(motion) loss_motion = Loss(pred, motion, mask) loss_vel = Loss.forward_vel(pred, motion, mask) loss = loss_motion + args.commit * loss_commit + args.loss_vel * loss_vel optimizer.zero_grad() loss.backward() optimizer.step() return loss_motion.item(), loss_commit.item(), perplexity.item(), loss_vel.item() ##### ------ warm-up ------- ##### avg = np.zeros(4) for nb_iter in range(1, args.warm_up_iter): optimizer, current_lr = update_lr_warm_up(optimizer, nb_iter, args.warm_up_iter, args.lr) motion, mask = next(train_loader_iter) motion, mask = motion.to(device).float(), mask.to(device).float() avg += step(motion, mask) if nb_iter % args.print_iter == 0: r, c, p, v = avg / args.print_iter logger.info(f"Warmup. Iter {nb_iter} : lr {current_lr:.6f} \t Commit. {c:.5f} " f"\t PPL. {p:.2f} \t Recons. {r:.5f} \t Vel. {v:.5f}") avg[:] = 0 ##### ---- Training ---- ##### best_hands, best_iter = 1e9, 0 avg[:] = 0 for nb_iter in range(1, args.total_iter + 1): motion, mask = next(train_loader_iter) motion, mask = motion.to(device).float(), mask.to(device).float() avg += step(motion, mask) scheduler.step() if nb_iter % args.print_iter == 0: r, c, p, v = avg / args.print_iter writer.add_scalar('./Train/Recons', r, nb_iter) writer.add_scalar('./Train/PPL', p, nb_iter) writer.add_scalar('./Train/Commit', c, nb_iter) logger.info(f"Train. Iter {nb_iter} : \t Commit. {c:.5f} \t PPL. {p:.2f} " f"\t Recons. {r:.5f} \t Vel. {v:.5f}") avg[:] = 0 if nb_iter % args.eval_iter == 0 or nb_iter == args.total_iter: recon, vel, mp, used = evaluate(net, val_loader, Loss, mean_t, std_t, device, groups=groups, n_kpts=layout.n_kpts) writer.add_scalar('./Val/Recons', recon, nb_iter) writer.add_scalar('./Val/MPJPE_hands', mp['hands'], nb_iter) writer.add_scalar('./Val/CodesUsed', used, nb_iter) logger.info(f"Eval. Iter {nb_iter} : val_recon {recon:.5f} val_vel {vel:.5f} " f"MPJPE all {mp['all']:.5f} body {mp['body']:.5f} " f"face {mp['face']:.5f} hands {mp['hands']:.5f} " f"| codes used {used}/{args.nb_code}") payload = {'net': net.state_dict(), 'args': vars(args), 'iter': nb_iter, 'val_recon': recon, 'val_mpjpe': mp, 'codes_used': used} torch.save(payload, os.path.join(args.out_dir, 'net_last.pth')) if mp['hands'] < best_hands: best_hands, best_iter = mp['hands'], nb_iter torch.save(payload, os.path.join(args.out_dir, 'net_best.pth')) logger.info(f" --> new best (hands MPJPE {best_hands:.5f})") logger.info(f"Done. best hands MPJPE {best_hands:.5f} @ iter {best_iter}") if __name__ == '__main__': main()