#!/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 `.npy` of int token ids per clip into --token-dir/. """ 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()