"""Argument parsers for the VSL (Full_TriVis) adaptation of T2M-GPT. Keeps the original T2M-GPT hyper-parameter names/defaults where they still apply and adds the ones this dataset needs (data dir, keypoint loss weights, the Vietnamese text encoder, token-sequence length). """ import argparse DIM_POSE = 256 # 128 DWPose keypoints x (x, y) def _common(parser): parser.add_argument('--data-dir', type=str, default='./dataset/VSL', help='output of prepare_vsl_data.py') parser.add_argument('--input-dim', type=int, default=DIM_POSE) parser.add_argument('--dataname', type=str, default='vsl') parser.add_argument('--seed', default=123, type=int) parser.add_argument('--device', type=str, default='cuda') parser.add_argument('--num-workers', type=int, default=8) parser.add_argument('--out-dir', type=str, default='output_vsl/') parser.add_argument('--exp-name', type=str, default='exp') parser.add_argument('--print-iter', default=100, type=int) parser.add_argument('--eval-iter', default=1000, type=int) ## vqvae arch (shared: stage 2 must rebuild the same net to decode) parser.add_argument("--code-dim", type=int, default=512) parser.add_argument("--nb-code", type=int, default=512) parser.add_argument("--mu", type=float, default=0.99) parser.add_argument("--down-t", type=int, default=2, help='2 -> 4x temporal downsample') parser.add_argument("--stride-t", type=int, default=2) parser.add_argument("--width", type=int, default=512) parser.add_argument("--depth", type=int, default=3) parser.add_argument("--dilation-growth-rate", type=int, default=3) parser.add_argument("--output-emb-width", type=int, default=512) parser.add_argument('--vq-act', type=str, default='relu', choices=['relu', 'silu', 'gelu']) parser.add_argument('--vq-norm', type=str, default=None) parser.add_argument("--quantizer", type=str, default='ema_reset', choices=['ema', 'orig', 'ema_reset', 'reset']) parser.add_argument('--beta', type=float, default=1.0) return parser def get_vq_args(): p = argparse.ArgumentParser(description='T2M-GPT stage 1 (pose VQ-VAE) on Full_TriVis', formatter_class=argparse.ArgumentDefaultsHelpFormatter) _common(p) p.add_argument('--batch-size', default=256, type=int) p.add_argument('--window-size', type=int, default=64, help='training clip length in frames') p.add_argument('--total-iter', default=50000, type=int) p.add_argument('--warm-up-iter', default=1000, type=int) p.add_argument('--lr', default=2e-4, type=float) p.add_argument('--lr-scheduler', default=[35000], nargs="+", type=int) p.add_argument('--gamma', default=0.05, type=float) p.add_argument('--weight-decay', default=0.0, type=float) p.add_argument("--commit", type=float, default=0.02) p.add_argument('--loss-vel', type=float, default=0.5) p.add_argument('--recons-loss', type=str, default='l1_smooth', choices=['l1', 'l2', 'l1_smooth']) ## keypoint weighting: hands carry the meaning in sign language p.add_argument('--w-body', type=float, default=1.0) p.add_argument('--w-face', type=float, default=0.5) p.add_argument('--w-hand', type=float, default=3.0) # --w-hand covers the whole 21-keypoint hand block. These two split it further: # --w-finger applies to every joint but the wrist, --w-fingertip to the five # tips only. Both default to None = fall back to --w-hand (original behaviour). p.add_argument('--w-finger', type=float, default=None) p.add_argument('--w-fingertip', type=float, default=None) p.add_argument('--val-windows', type=int, default=2000, help='cap on deterministic val windows (0 = all)') p.add_argument("--resume-pth", type=str, default=None) return p.parse_args() def get_trans_args(): p = argparse.ArgumentParser(description='T2M-GPT stage 2 (text->pose-token GPT) on Full_TriVis', formatter_class=argparse.ArgumentDefaultsHelpFormatter) _common(p) p.add_argument('--batch-size', default=64, type=int) p.add_argument('--total-iter', default=100000, type=int) p.add_argument('--lr', default=1e-4, type=float) p.add_argument('--lr-scheduler', default=[60000], nargs="+", type=int) p.add_argument('--gamma', default=0.05, type=float) p.add_argument('--weight-decay', default=1e-6, type=float) p.add_argument('--optimizer', type=str, default='adamw', choices=['adam', 'adamw']) p.add_argument('--decay-option', type=str, default='all', choices=['all', 'noVQ']) ## gpt arch p.add_argument('--token-dir', type=str, default='./dataset/VSL/tokens', help='per-clip VQ token .npy from tokenize_vsl.py') p.add_argument("--max-tokens", type=int, default=128, help='max pose-token sequence length (clip frames / 2**down_t)') # The paper's HumanML3D config is embed_dim_gpt 1024 / n_head_gpt 16 (~227M params). # Full_TriVis has 19k gloss-pose pairs (9.5k unique gloss) vs HumanML3D's ~69k # text-motion pairs, so the narrower 512/8 (~58M) is the default here; pass the # paper's values explicitly to reproduce its capacity. p.add_argument("--embed-dim-gpt", type=int, default=512) p.add_argument("--text-model", type=str, default='vinai/phobert-base-v2', help='frozen Vietnamese text encoder replacing CLIP') p.add_argument("--text-field", type=str, default='gloss', choices=['gloss', 'sentence']) p.add_argument("--num-layers", type=int, default=9) p.add_argument("--n-head-gpt", type=int, default=8) p.add_argument("--ff-rate", type=int, default=4) p.add_argument("--drop-out-rate", type=float, default=0.1) p.add_argument("--pkeep", type=float, default=0.5, help='token-corruption keep prob (T2M-GPT corruption trick)') p.add_argument("--resume-pth", type=str, required=True, help='trained stage-1 VQ-VAE') p.add_argument("--resume-trans", type=str, default=None) p.add_argument("--init-trans", type=str, default=None, help='warm-start stage 2 from a checkpoint trained on another corpus ' '(shape-compatible tensors only); ignored if --resume-trans is set') p.add_argument('--eval-samples', type=int, default=300, help='clips used for the generation-based val metrics') # 0 = off (previous behaviour). Set e.g. 10000 to keep net_iter10000.pth etc, so a # long run can be scored at FIXED iterations afterwards. net_last is overwritten and # net_best is selected on the noisy 80-clip metric, so neither supports a # training-length study. ~230 MB per snapshot. p.add_argument('--save-every', type=int, default=0, help='also snapshot net_iter.pth every N iters (0 = off)') return p.parse_args()