| |
| """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() |
|
|
| |
| 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)) |
|
|
| |
| 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() |
| |
| 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}') |
|
|
| |
| 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() |
|
|
| |
| 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 |
|
|
| |
| 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() |
|
|