t2m-gpt-vsl-code / options /option_vsl.py
Tri1's picture
T2M-GPT VSL adaptation: Python sources only (82 files, no checkpoints or data)
8e5456b verified
Raw
History Blame Contribute Delete
6.96 kB
"""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()