| |
| """Save an UNTRAINED stage-2 checkpoint, as the iter-0 control for eval_vsl.py. |
| |
| Without this baseline, a flat generation-metric curve is ambiguous: it could mean |
| "training converged in the first few thousand iters" or "the transformer barely |
| contributes and the number reflects the pose prior". Scoring a randomly |
| initialized GPT through the *trained* VQ-VAE decoder separates the two. |
| |
| Reuses the trained checkpoint's args so the architecture and eval path are |
| identical -- only the transformer weights are random. |
| """ |
| import argparse |
|
|
| import torch |
|
|
| import models.t2m_trans as trans |
| from models.text_encoder_vi import ViTextEncoder |
|
|
|
|
| def main(): |
| ap = argparse.ArgumentParser() |
| ap.add_argument('--like', required=True, help='trained stage-2 ckpt to copy args from') |
| ap.add_argument('--out', required=True) |
| ap.add_argument('--seed', type=int, default=1234) |
| args = ap.parse_args() |
|
|
| ref = torch.load(args.like, map_location='cpu') |
| gargs = argparse.Namespace(**ref['args']) |
| torch.manual_seed(args.seed) |
|
|
| |
| dim = ViTextEncoder(gargs.text_model, device='cpu').dim |
| net = trans.Text2Motion_Transformer( |
| num_vq=gargs.nb_code, embed_dim=gargs.embed_dim_gpt, clip_dim=dim, |
| block_size=gargs.max_tokens + 1, num_layers=gargs.num_layers, |
| n_head=gargs.n_head_gpt, drop_out_rate=gargs.drop_out_rate, fc_rate=gargs.ff_rate) |
|
|
| torch.save({'trans': net.state_dict(), 'args': ref['args'], 'iter': 0, |
| 'val_loss': float('nan'), 'val_acc': float('nan'), |
| 'note': 'UNTRAINED random init - iter-0 control'}, args.out) |
| print(f'wrote untrained control checkpoint -> {args.out}') |
|
|
|
|
| if __name__ == '__main__': |
| main() |
|
|