MCPcopy Create free account
hub / github.com/THUDM/P-tuning / construct_generation_args

Function construct_generation_args

LAMA/cli.py:31–76  ·  view source on GitHub ↗
()

Source from the content-addressed store, hash-verified

29
30
31def construct_generation_args():
32 parser = argparse.ArgumentParser()
33
34 # pre-parsing args
35 parser.add_argument("--relation_id", type=str, default="P1001")
36 parser.add_argument("--model_name", type=str, default='megatron_11b', choices=SUPPORT_MODELS)
37 parser.add_argument("--pseudo_token", type=str, default='[PROMPT]')
38
39 parser.add_argument("--t5_shard", type=int, default=0)
40 parser.add_argument("--mid", type=int, default=0)
41 parser.add_argument("--template", type=str, default="(3, 3, 3)")
42 parser.add_argument("--early_stop", type=int, default=20)
43
44 parser.add_argument("--lr", type=float, default=1e-5)
45 parser.add_argument("--seed", type=int, default=34, help="random seed for initialization")
46 parser.add_argument("--decay_rate", type=float, default=0.98)
47 parser.add_argument("--weight_decay", type=float, default=0.0005)
48 parser.add_argument("--no_cuda", action="store_true", help="Avoid using CUDA when available")
49
50 # lama configuration
51 parser.add_argument("--only_evaluate", type=bool, default=False)
52 parser.add_argument("--use_original_template", type=bool, default=False)
53 parser.add_argument("--use_lm_finetune", type=bool, default=False)
54
55 parser.add_argument("--vocab_strategy", type=str, default="shared", choices=['original', 'shared', 'lama'])
56 parser.add_argument("--lstm_dropout", type=float, default=0.0)
57
58 # directories
59 parser.add_argument("--data_dir", type=str, default=join(abspath(dirname(__file__)), '../data/LAMA'))
60 parser.add_argument("--out_dir", type=str, default=join(abspath(dirname(__file__)), '../out/LAMA'))
61 # MegatronLM 11B
62 parser.add_argument("--checkpoint_dir", type=str, default=join(abspath(dirname(__file__)), '../checkpoints'))
63
64 args = parser.parse_args()
65
66 # post-parsing args
67
68 args.device = torch.device("cuda" if torch.cuda.is_available() and not args.no_cuda else "cpu")
69 args.n_gpu = 0 if args.no_cuda else torch.cuda.device_count()
70 args.template = eval(args.template) if type(args.template) is not tuple else args.template
71
72 assert type(args.template) is tuple
73
74 set_seed(args)
75
76 return args
77
78
79class Trainer(object):

Callers 1

mainFunction · 0.85

Calls 1

set_seedFunction · 0.70

Tested by

no test coverage detected