File size: 6,955 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
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
"""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<N>.pth every N iters (0 = off)')
    return p.parse_args()