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