Repository navigation
Expand file tree
/
Copy pathtrain.py
More file actions
139 lines (104 loc) · 4.16 KB
/
Copy pathtrain.py
File metadata and controls
139 lines (104 loc) · 4.16 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
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
import time
import os
import sys
import torch
import tqdm
import sentencepiece as spm
import random
from keras_preprocessing.sequence import pad_sequences
from transformer_lm import TransformerLM
if len(sys.argv) != 3:
print("Expected 2 arguments (path to model, path to dataset)")
exit(1)
model_path = os.path.abspath(sys.argv[1])
dataset_path = os.path.abspath(sys.argv[2])
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
def load(filename):
file = open(filename, "r", encoding="utf-8")
examples = []
for line in file:
line = line[:-len("\n")]
examples.append(line)
return examples
tokenizer = spm.SentencePieceProcessor()
tokenizer.Load(os.path.join(dataset_path, "code_spm.model"))
tokenizer.SetEncodeExtraOptions("bos:eos")
print("Loading dataset...")
leclair_train = load(os.path.join(dataset_path, "train_codes.txt"))
leclair_val = load(os.path.join(dataset_path, "val_codes.txt"))
print("Creating model...")
model = TransformerLM.from_description(os.path.join(model_path, "model_description.json")).to(device)
criterion = torch.nn.CrossEntropyLoss(ignore_index=0)
optimizer = torch.optim.Adam(model.parameters())
scheduler = torch.optim.lr_scheduler.StepLR(optimizer, 1, gamma=0.95)
batch_size = 80
random.seed()
def pad_and_split(batch, max_item_len):
padded_batch = pad_sequences(batch, maxlen=max_item_len, padding='post', truncating='post',
value=0, dtype='int64')
padded_batch = torch.tensor(padded_batch).to(device)
context = padded_batch[:, :-1]
target = padded_batch[:, 1:]
return context, target
def batch_dataset(dataset):
random.shuffle(dataset)
num_batches = (len(dataset) - 1) // batch_size
max_item_len = model.input_length + 1
for batch_num in tqdm.trange(num_batches):
batch = dataset[batch_num * batch_size: batch_num * batch_size + batch_size]
tokenized_batch = []
for item in batch:
tokenized = tokenizer.SampleEncodeAsIds(item, -1, 0.2)
if len(tokenized) <= max_item_len:
tokenized_batch.append(tokenized)
yield pad_and_split(tokenized_batch, max_item_len)
def train():
model.train()
total_loss = 0.
start_time = time.time()
counter = 0
for context, target in batch_dataset(leclair_train):
counter += 1
optimizer.zero_grad()
output = model(context)
loss = criterion(output.view(-1, model.vocab_size), target.flatten())
loss.backward()
torch.nn.utils.clip_grad_norm_(model.parameters(), 0.5)
optimizer.step()
total_loss += loss.item()
log_interval = 200
if counter % log_interval == 0 and counter > 0:
cur_loss = total_loss / counter
elapsed = time.time() - start_time
print("Batch %d, current train loss: %.4f, time elapsed: %.1f"
% (counter, cur_loss, elapsed))
def evaluate(data_source):
model.eval()
total_loss = 0.
batch_counter = 0
with torch.no_grad():
for context, target in batch_dataset(data_source):
output = model(context)
total_loss += criterion(output.view(-1, model.vocab_size), target.flatten()).item()
batch_counter += 1
return total_loss / batch_counter
model_save_path = os.path.join(model_path, "trained_model")
if os.path.exists(model_save_path):
print("Loading and evaluating existing model...")
model.load_state_dict(torch.load(model_save_path, map_location=device))
best_val_loss = evaluate(leclair_val)
print("Initial validation loss: %.4f" % best_val_loss)
else:
print("No existing model, starting training from scratch")
best_val_loss = float("inf")
epochs = 15
for epoch in range(1, epochs + 1):
print("\nStarting epoch %d of %d" % (epoch, epochs))
epoch_start_time = time.time()
train()
val_loss = evaluate(leclair_val)
print("Finished epoch %d of %d, val loss: %.4f" % (epoch, epochs, val_loss))
if val_loss < best_val_loss:
print("Val loss improved from %.4f, saving checkpoint" % best_val_loss)
best_val_loss = val_loss
torch.save(model.state_dict(), model_save_path)