| 1 | import argparse |
| 2 | |
| 3 | def parse_args(): |
| 4 | parser = argparse.ArgumentParser(description="Run DCCF.") |
| 5 | parser.add_argument('--data_path', nargs='?', default='data/', help='Input data path.') |
| 6 | parser.add_argument('--seed', type=int, default=2022, help='random seed') |
| 7 | parser.add_argument('--dataset', nargs='?', default='gowalla', help='Choose a dataset from {gowalla, amazon, tmall}') |
| 8 | parser.add_argument('--verbose', type=int, default=1, help='Interval of evaluation.') |
| 9 | parser.add_argument('--save_model', type=bool, default=False, help='Whether to save') |
| 10 | parser.add_argument('--epoch', type=int, default=100, help='Number of epochs') |
| 11 | parser.add_argument('--embed_size', type=int, default=32, help='Embedding size.') |
| 12 | parser.add_argument('--n_batch', type=int, default=40, help='Number of mini-batches') |
| 13 | parser.add_argument('--batch_size', type=int, default=10240, help='batch size') |
| 14 | parser.add_argument('--train_num', type=int, default=10000, help='Number of training instances per epoch') |
| 15 | parser.add_argument('--sample_num', type=int, default=40, help='Number of pos/neg samples for each instance') |
| 16 | parser.add_argument('--lr', type=float, default=0.001, help='Learning rate.') |
| 17 | parser.add_argument('--emb_reg', type=float, default=2.5e-5, help='Regularizations.') |
| 18 | parser.add_argument('--cen_reg', type=float, default=5e-3, help='Regularizations.') |
| 19 | parser.add_argument('--ssl_reg', type=float, default=1e-1, help='Reg weight for ssl loss') |
| 20 | parser.add_argument('--n_layers', type=int, default=2, help='Layer numbers.') |
| 21 | parser.add_argument('--n_intents', type=int, default=128, help='Number of latent intents') |
| 22 | parser.add_argument('--temp', type=float, default=1, help='temperature in ssl loss') |
| 23 | parser.add_argument('--show_step', type=int, default=1, help='Test every show_step epochs.') |
| 24 | parser.add_argument('--Ks', nargs='?', default='[20, 40]', help='Metrics scale') |
| 25 | |
| 26 | return parser.parse_args() |