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