MCPcopy Create free account
hub / github.com/THUNLP-MT/MEAN / parse

Function parse

train.py:14–46  ·  view source on GitHub ↗
()

Source from the content-addressed store, hash-verified

12
13
14def 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
49def prepare_efficient_mc_att(args):

Callers 1

train.pyFile · 0.70

Calls

no outgoing calls

Tested by

no test coverage detected