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()