| 12 | from trainer import Trainer |
| 13 | |
| 14 | def parse_args(): |
| 15 | parser = argparse.ArgumentParser(description="Index") |
| 16 | |
| 17 | parser.add_argument('--lr', type=float, default=1e-3, help='learning rate') |
| 18 | parser.add_argument('--epochs', type=int, default=5000, help='number of epochs') |
| 19 | parser.add_argument('--batch_size', type=int, default=2048, help='batch size') |
| 20 | parser.add_argument('--num_workers', type=int, default=4, ) |
| 21 | parser.add_argument('--eval_step', type=int, default=50, help='eval step') |
| 22 | parser.add_argument('--learner', type=str, default="AdamW", help='optimizer') |
| 23 | parser.add_argument('--lr_scheduler_type', type=str, default="constant", help='scheduler') |
| 24 | parser.add_argument('--warmup_epochs', type=int, default=50, help='warmup epochs') |
| 25 | parser.add_argument("--data_path", type=str, |
| 26 | default="../data/Games/Games.emb-llama-td.npy", |
| 27 | help="Input data path.") |
| 28 | |
| 29 | parser.add_argument("--weight_decay", type=float, default=0.0, help='l2 regularization weight') |
| 30 | parser.add_argument("--dropout_prob", type=float, default=0.0, help="dropout ratio") |
| 31 | parser.add_argument("--bn", type=bool, default=False, help="use bn or not") |
| 32 | parser.add_argument("--loss_type", type=str, default="mse", help="loss_type") |
| 33 | parser.add_argument("--kmeans_init", type=bool, default=True, help="use kmeans_init or not") |
| 34 | parser.add_argument("--kmeans_iters", type=int, default=100, help="max kmeans iters") |
| 35 | parser.add_argument('--sk_epsilons', type=float, nargs='+', default=[0.0, 0.0, 0.0], help="sinkhorn epsilons") |
| 36 | parser.add_argument("--sk_iters", type=int, default=50, help="max sinkhorn iters") |
| 37 | |
| 38 | parser.add_argument("--device", type=str, default="cuda:0", help="gpu or cpu") |
| 39 | |
| 40 | parser.add_argument('--num_emb_list', type=int, nargs='+', default=[256,256,256], help='emb num of every vq') |
| 41 | parser.add_argument('--e_dim', type=int, default=32, help='vq codebook embedding size') |
| 42 | parser.add_argument('--quant_loss_weight', type=float, default=1.0, help='vq quantion loss weight') |
| 43 | parser.add_argument("--beta", type=float, default=0.25, help="Beta for commitment loss") |
| 44 | parser.add_argument('--layers', type=int, nargs='+', default=[2048,1024,512,256,128,64], help='hidden sizes of every layer') |
| 45 | |
| 46 | parser.add_argument('--save_limit', type=int, default=5) |
| 47 | parser.add_argument("--ckpt_dir", type=str, default="", help="output directory for model") |
| 48 | |
| 49 | return parser.parse_args() |
| 50 | |
| 51 | |
| 52 | if __name__ == '__main__': |