-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathdual_api_run.py
More file actions
58 lines (54 loc) · 2.24 KB
/
Copy pathdual_api_run.py
File metadata and controls
58 lines (54 loc) · 2.24 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
from models.dual_model_api import OrthrusDualAPI
from utils import set_logger
from datetime import datetime
import pytz
import logging
import argparse
import os
if __name__ == '__main__':
parser = argparse.ArgumentParser()
parser.add_argument("--max_length", type=int, required=True,
help="max length")
parser.add_argument("--batch_size", type=int, required=True,
help="training batch size")
parser.add_argument("--epoch", type=int, required=True,
help="number of epochs")
parser.add_argument("--fuse", type=bool, default=False,
help="fuse or not")
parser.add_argument("--norm", type=bool, default=False,
help="normalize or not")
parser.add_argument("--output_dir", type=str, required=True,
help="output directory")
parser.add_argument("--load_model_path", default=None,
help="load model from")
parser.add_argument("--training_file", type=str, required=True, default='../data/061222_training.csv')
parser.add_argument("--valid_file", type=str, required=True, default='../data/061222_valid.csv')
args = parser.parse_args()
os.makedirs('./log', exist_ok=True)
set_logger('./log/dual_{}.log'.format(datetime.now(pytz.timezone('Asia/Singapore'))))
logging.info(args)
logging.info('Training Dual Model - Annotation + SO Title + SO API')
# Load the fine-tuned model
model = OrthrusDualAPI(codebert_path = 'microsoft/codebert-base',
decoder_layers = 6,
fix_encoder = False,
beam_size = 5,
max_source_length = args.max_length,
max_target_length = args.max_length,
load_model_path = args.load_model_path,
l2_norm = args.norm,
fusion = args.fuse
)
# train model
model.train(
# train_filename ='../data/train.csv',
train_filename = args.training_file,
train_batch_size = args.batch_size,
num_train_epochs = args.epoch,
learning_rate = 5e-5,
do_eval = True,
# dev_filename ='../data/valid.csv',
dev_filename = args.valid_file,
eval_batch_size = 64,
output_dir = args.output_dir
)