| 23 | |
| 24 | |
| 25 | def parse_args(): |
| 26 | parser = argparse.ArgumentParser(description="Run supervised GRU.") |
| 27 | |
| 28 | parser.add_argument('--epoch', type=int, default=500, |
| 29 | help='Number of max epochs.') |
| 30 | parser.add_argument('--data', nargs='?', default='Goodreads_5', |
| 31 | help='Toys_and_Games, Goodreads, Industrial_and_Scientific, CDs_and_Vinyl') |
| 32 | # parser.add_argument('--pretrain', type=int, default=1, |
| 33 | # help='flag for pretrain. 1: initialize from pretrain; 0: randomly initialize; -1: save the model to pretrain file') |
| 34 | parser.add_argument('--batch_size', type=int, default=1024, |
| 35 | help='Batch size.') |
| 36 | parser.add_argument('--hidden_factor', type=int, default=32, |
| 37 | help='Number of hidden factors, i.e., embedding size.') |
| 38 | parser.add_argument('--num_filters', type=int, default=16, |
| 39 | help='num_filters') |
| 40 | parser.add_argument('--filter_sizes', nargs='?', default='[2,3,4]', |
| 41 | help='Specify the filter_size') |
| 42 | parser.add_argument('--r_click', type=float, default=0.2, |
| 43 | help='reward for the click behavior.') |
| 44 | parser.add_argument('--r_buy', type=float, default=1.0, |
| 45 | help='reward for the purchase behavior.') |
| 46 | parser.add_argument('--lr', type=float, default=0.001, |
| 47 | help='Learning rate.') |
| 48 | parser.add_argument('--save_flag', type=int, default=1, |
| 49 | help='0: Disable model saver, 1: Activate model saver') |
| 50 | parser.add_argument('--cuda', type=int, default=1, |
| 51 | help='cuda device.') |
| 52 | parser.add_argument('--l2_decay', type=float, default=1e-5, |
| 53 | help='l2 loss reg coef.') |
| 54 | parser.add_argument('--alpha', type=float, default=0, |
| 55 | help='dro alpha.') |
| 56 | parser.add_argument('--beta', type=float, default=1.0, |
| 57 | help='for robust radius') |
| 58 | parser.add_argument("--model", type=str, default="SASRec", |
| 59 | help='the model name, GRU, Caser, SASRec') |
| 60 | parser.add_argument('--dropout_rate', type=float, default=0.3, |
| 61 | help='dropout ') |
| 62 | parser.add_argument('--descri', type=str, default='', |
| 63 | help='description of the work.') |
| 64 | parser.add_argument("--early_stop", type=int, default=20, |
| 65 | help='the epoch for early stop') |
| 66 | parser.add_argument("--eval_num", type=int, default=1, |
| 67 | help='evaluate every eval_num epoch' ) |
| 68 | parser.add_argument("--seed", type=int, default=1, |
| 69 | help="the random seed") |
| 70 | parser.add_argument("--result_json_path", type=str, default="./result_temp/temp.json") |
| 71 | parser.add_argument("--sample_num", type=int, default = 65536) |
| 72 | parser.add_argument("--debug", type=bool, default=False) |
| 73 | parser.add_argument("--loss_type", type=str, default="bce") |
| 74 | return parser.parse_args() |
| 75 | |
| 76 | def setup_seed(seed): |
| 77 | torch.manual_seed(seed) |