File size: 3,247 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
#!/usr/bin/env python3
"""Encode every Full_TriVis clip to a VQ token sequence with a trained stage-1 net.

Equivalent to the `##### ---- get code ---- #####` block inside T2M-GPT's
train_t2m_trans.py, pulled out into its own script so stage 2 can restart
without re-encoding, and extended to all three splits (the original only
tokenizes train).

Writes one `<clip_name>.npy` of int token ids per clip into --token-dir/<split>.
"""
import argparse
import json
import os

import numpy as np
import torch
from tqdm import tqdm

import models.vqvae as vqvae
from dataset import dataset_vsl


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument('--data-dir', default='./dataset/VSL')
    ap.add_argument('--resume-pth', required=True, help='trained stage-1 checkpoint')
    ap.add_argument('--token-dir', default='./dataset/VSL/tokens')
    ap.add_argument('--splits', nargs='+', default=['train', 'val', 'test'])
    ap.add_argument('--device', default='cuda')
    ap.add_argument('--max-frames', type=int, default=512,
                    help='truncate very long clips (0 = keep full length)')
    args = ap.parse_args()

    device = torch.device(args.device)
    ckpt = torch.load(args.resume_pth, map_location='cpu')
    targs = argparse.Namespace(**ckpt['args'])
    print(f"stage-1 ckpt: iter {ckpt.get('iter')} val_recon {ckpt.get('val_recon')} "
          f"mpjpe {ckpt.get('val_mpjpe')} codes_used {ckpt.get('codes_used')}")

    net = vqvae.HumanVQVAE(targs, targs.nb_code, targs.code_dim, targs.output_emb_width,
                           targs.down_t, targs.stride_t, targs.width, targs.depth,
                           targs.dilation_growth_rate, targs.vq_act, targs.vq_norm,
                           input_dim=targs.input_dim)
    net.load_state_dict(ckpt['net'], strict=True)
    net.eval().to(device)

    unit_length = 2 ** targs.down_t
    stats = {}
    for split in args.splits:
        out_dir = os.path.join(args.token_dir, split)
        os.makedirs(out_dir, exist_ok=True)
        ds = dataset_vsl.VSLTokenizeDataset(args.data_dir, split, unit_length=unit_length,
                                           max_frames=args.max_frames)
        loader = torch.utils.data.DataLoader(ds, batch_size=1, shuffle=False, num_workers=6)

        lens, used = [], set()
        with torch.no_grad():
            for motion, name, _ in tqdm(loader, desc=f'tokenize {split}'):
                motion = motion.to(device).float()
                tokens = net.encode(motion)[0].cpu().numpy().astype(np.int32)
                np.save(os.path.join(out_dir, name[0] + '.npy'), tokens)
                lens.append(len(tokens))
                used.update(tokens.tolist())

        L = np.array(lens)
        stats[split] = {'clips': int(len(L)), 'tokens_median': int(np.median(L)),
                        'tokens_p95': int(np.percentile(L, 95)), 'tokens_max': int(L.max()),
                        'codes_used': len(used)}
        print(split, stats[split])

    with open(os.path.join(args.token_dir, 'stats.json'), 'w') as f:
        json.dump({'stage1_ckpt': args.resume_pth, 'unit_length': unit_length,
                   'nb_code': targs.nb_code, 'splits': stats}, f, indent=2)


if __name__ == '__main__':
    main()