-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathtrain.py
More file actions
146 lines (120 loc) · 5.93 KB
/
Copy pathtrain.py
File metadata and controls
146 lines (120 loc) · 5.93 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
140
141
142
143
144
145
146
import argparse
import os
import numpy as np
import torch
import torch.nn as nn
from torch.nn.utils.rnn import pack_padded_sequence
from torch.utils.data import DataLoader
from file_path_manager import FilePathManager
from inceptionv3_extractor import InceptionV3Extractor
from misc.coco_dataset import CocoDataset
from misc.corpus import Corpus
from model import m_RNN
from vgg16_extractor import Vgg16Extractor
def handle(x, cuda):
if cuda and torch.cuda.is_available():
x = x.cuda()
return x
def main(args):
use_cuda = True
# Create model directory
if not os.path.exists(args.model_path):
os.makedirs(args.model_path)
if args.image_regions == 64:
# Inception V3 feature extractor
extractor = InceptionV3Extractor(use_gpu=use_cuda, transform=False)
else:
# VGG feature extractor
extractor = Vgg16Extractor(use_gpu=use_cuda, transform=False, regions_count=args.image_regions)
# Load vocabulary wrapper.
corpus = Corpus.load(FilePathManager.resolve(args.corpus_path))
print(corpus.word_from_index(0))
dataset = CocoDataset(corpus, root=args.image_dir, annFile=args.caption_path, transform=extractor.tf_image)
# Build data loader
dataloader = DataLoader(dataset, batch_size=args.batch_size, shuffle=True, num_workers=args.num_workers,
pin_memory=use_cuda)
# Build the models
model = m_RNN(use_cuda=use_cuda,
image_regions=extractor.regions_count,
regions_features=extractor.regions_features_size,
features_size=extractor.features_size)
if torch.cuda.is_available():
model.cuda()
# Loss and Optimizer
# criterion = nn.CrossEntropyLoss(ignore_index=corpus.word_index(corpus.PAD))
criterion = nn.CrossEntropyLoss()
params = list(model.parameters())
optimizer = torch.optim.Adam(params, lr=args.lr, weight_decay=args.w_decay)
# Continue Training
if args.pre_trained_epoch != 0:
model.load_state_dict(torch.load(f"{args.model_path}model-{args.pre_trained_epoch}.pkl"))
optimizer.load_state_dict(torch.load(f"{args.model_path}optimizer-{args.pre_trained_epoch}.pkl"))
epoch_loss = 0
log = ''
print('enable Logging')
# Train the Models
total_step = len(dataloader)
for epoch in range(args.pre_trained_epoch, args.pre_trained_epoch + 1):
for i, (images, inputs, targets, _) in enumerate(dataloader):
images = images.cuda()
images_features, images_regions = extractor.forward(images)
for k in range(inputs.shape[1]):
# Set mini-batch dataset
images = handle(images, cuda=use_cuda)
input = handle(inputs[:, k, :-1], cuda=use_cuda)
target = handle(targets[:, k, 1:], cuda=use_cuda)
input = pack_padded_sequence(input, [17] * input.shape[0], True)[0]
target = pack_padded_sequence(target, [17] * target.shape[0], True)[0]
# Forward, Backward and Optimize
model.zero_grad()
# make update
output = model(images_features, images_regions, input)
loss = criterion(output, target)
loss.backward()
optimizer.step()
epoch_loss += loss.item()
# Print log info
if i % args.log_step == 0:
print('Epoch [%d/%d], Step [%d/%d], Loss: %.4f, Perplexity: %5.4f'
% (epoch, args.pre_trained_epoch + 1, i, total_step,
loss.item(), np.exp(loss.item())))
log += 'Epoch [%d/%d], Step [%d/%d], Loss: %.4f, Perplexity: %5.4f \n' \
% (epoch, args.pre_trained_epoch + 1, i, total_step,
loss.item(), np.exp(loss.item()))
# Save the models
torch.save(model.state_dict(),
os.path.join(args.model_path,
'model-%d.pkl' % (epoch + 1)))
torch.save(optimizer.state_dict(),
os.path.join(args.model_path,
'optimizer-%d.pkl' % (epoch + 1)))
print(f'epoch {epoch+1} loss is : {epoch_loss}')
log += f'epoch {epoch+1} loss is : {epoch_loss}\n'
with open(os.path.join(args.model_path, 'log-%d.txt' % (epoch + 1)), "w") as text_file:
text_file.write(log)
if __name__ == '__main__':
parser = argparse.ArgumentParser()
parser.add_argument('--model_path', type=str, default='./models/',
help='path for saving trained models')
parser.add_argument('--pre_trained_epoch', type=int, default=0,
help='path for saved trained models')
parser.add_argument('--corpus_path', type=str, default='data/corpus.pkl',
help='path for vocabulary wrapper')
parser.add_argument('--caption_path', type=str,
default='D:/Datasets/mscoco/test/captions_train2017.json',
help='path for train annotation json file')
parser.add_argument('--image_dir', type=str, default='D:/Datasets/mscoco/test/images',
help='directory for resized images')
parser.add_argument('--log_step', type=int, default=1,
help='step size for printing log info')
parser.add_argument('--save_step', type=int, default=2,
help='step size for saving trained models')
parser.add_argument('--image_regions', type=int, default=64,
help='number of image regions to be extracted (49 or 196) 64 for inception_v3')
parser.add_argument('--num_epochs', type=int, default=1)
parser.add_argument('--batch_size', type=int, default=5)
parser.add_argument('--num_workers', type=int, default=0)
parser.add_argument('--lr', type=float, default=0.0001)
parser.add_argument('--w_decay', type=float, default=0)
args = parser.parse_args()
main(args)