| 12 | |
| 13 | |
| 14 | def parse(): |
| 15 | parser = argparse.ArgumentParser(description='training') |
| 16 | parser.add_argument('--train_set', type=str, required=True, help='path to train set') |
| 17 | parser.add_argument('--valid_set', type=str, required=True, help='path to valid set') |
| 18 | |
| 19 | # training related |
| 20 | parser.add_argument('--lr', type=float, default=1e-3, help='learning rate') |
| 21 | parser.add_argument('--max_epoch', type=int, default=10, help='max training epoch') |
| 22 | parser.add_argument('--grad_clip', type=float, default=1.0, help='clip gradients with too big norm') |
| 23 | parser.add_argument('--save_dir', type=str, required=True, help='directory to save model and logs') |
| 24 | parser.add_argument('--batch_size', type=int, required=True, help='batch size') |
| 25 | parser.add_argument('--shuffle', action='store_true', help='shuffle data') |
| 26 | parser.add_argument('--num_workers', type=int, default=4) |
| 27 | parser.add_argument('--mode', type=str, default='111', help='H/L/Antigen, 1 for include, 0 for exclude') |
| 28 | parser.add_argument('--seed', type=int, default=42, help='Seed to use in training') |
| 29 | parser.add_argument('--early_stop', action='store_true', help='Whether to use early stop') |
| 30 | |
| 31 | # device |
| 32 | parser.add_argument('--gpus', type=int, nargs='+', required=True, help='gpu to use, -1 for cpu') |
| 33 | parser.add_argument("--local_rank", type=int, default=-1, |
| 34 | help="Local rank. Necessary for using the torch.distributed.launch utility.") |
| 35 | |
| 36 | ## shared |
| 37 | parser.add_argument('--cdr_type', type=str, default='3', help='type of cdr') |
| 38 | ## for Multi-Channel Attetion model |
| 39 | parser.add_argument('--embed_size', type=int, default=64, help='embed size of amino acids') |
| 40 | parser.add_argument('--hidden_size', type=int, default=128, help='hidden size') |
| 41 | parser.add_argument('--n_layers', type=int, default=3, help='number of layers') |
| 42 | parser.add_argument('--alpha', type=float, default=0.05, help='scale mse loss of coordinates') |
| 43 | parser.add_argument('--anneal_base', type=float, default=1, help='Exponential lr decay, 1 for not decay') |
| 44 | ## for efficient version |
| 45 | parser.add_argument('--n_iter', type=int, default=5, help='Number of iterations') |
| 46 | return parser.parse_args() |
| 47 | |
| 48 | |
| 49 | def prepare_efficient_mc_att(args): |