MCPcopy Create free account
hub / github.com/AkaliKong/MiniOneRec / parse_args

Function parse_args

rq/rqvae.py:14–49  ·  view source on GitHub ↗
()

Source from the content-addressed store, hash-verified

12from trainer import Trainer
13
14def 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
52if __name__ == '__main__':

Callers 1

rqvae.pyFile · 0.70

Calls

no outgoing calls

Tested by

no test coverage detected