t2m-gpt-vsl-code / train_vq_vsl.py
Tri1's picture
T2M-GPT VSL adaptation: Python sources only (82 files, no checkpoints or data)
8e5456b verified
Raw
History Blame Contribute Delete
8.57 kB
#!/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()