-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathrun.py
More file actions
86 lines (74 loc) · 2.72 KB
/
Copy pathrun.py
File metadata and controls
86 lines (74 loc) · 2.72 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
# pylint: disable=import-error
import warnings
warnings.filterwarnings("ignore")
import argparse
import json
import random
import torch
import numpy as np
from online_train import train
from models import load_model
def set_seed(seed):
random.seed(seed)
np.random.seed(seed)
torch.manual_seed(seed)
if torch.cuda.is_available():
torch.cuda.manual_seed_all(seed)
if __name__ == '__main__':
parser = argparse.ArgumentParser()
parser.add_argument('--config', default=None, type=str)
arg_ = parser.parse_args()
if arg_.config is None:
raise NameError("Include a config file in the argument please.")
# Getting configurations
with open(arg_.config) as config_file:
hparam = json.load(config_file)
hparam = argparse.Namespace(**hparam)
args_dict = dict(
output_dir=hparam.__dict__.get('output_dir'),
model_editing_config=hparam.__dict__.get('model_editing_config'),
dataset=hparam.dataset,
dataset_version=hparam.dataset_version,
model_name_or_path=hparam.model,
method=hparam.method,
coreset=hparam.coreset,
coreset_ratio=hparam.coreset_ratio,
freeze_level=hparam.__dict__.get('freeze_level'),
max_input_length=hparam.input_length,
max_output_length=hparam.output_length,
freeze_encoder=False,
freeze_embeds=False,
learning_rate=hparam.__dict__.get('learning_rate'),
weight_decay=0.0,
adam_epsilon=1e-8,
warmup_steps=0,
train_batch_size=hparam.__dict__.get('train_batch_size'),
eval_batch_size=hparam.__dict__.get('eval_batch_size'),
num_train_epochs=hparam.__dict__.get('num_train_epochs'),
gradient_accumulation_steps=hparam.__dict__.get(
'gradient_accumulation_steps'),
n_gpu=hparam.ngpu,
repeat_num=hparam.repeat_num,
num_workers=4 * hparam.ngpu,
use_lr_scheduling=hparam.__dict__.get('use_lr_scheduling'),
val_check_interval=1.0,
use_deepspeed=hparam.__dict__.get('use_deepspeed'),
max_grad_norm=0.5,
seed=42,
check_validation_only=hparam.check_validation,
checkpoint_path=hparam.__dict__.get('checkpoint_path'),
output_log=hparam.__dict__.get('output_log'),
red_flag=hparam.red_flag,
alpha=hparam.__dict__.get('alpha'),
temperature=hparam.__dict__.get('temperature'),
distil_epoch=hparam.__dict__.get('distil_epoch')
)
args = argparse.Namespace(**args_dict)
if 't5' in args.model_name_or_path:
Model = load_model('T5')
elif 'llama' in args.model_name_or_path:
Model = load_model('Llama')
else:
Model = load_model('GPT2')
set_seed(42)
train(args, Model)