File size: 8,569 Bytes
8e5456b | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 | #!/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()
|