-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathtrain.py
More file actions
72 lines (59 loc) · 3.3 KB
/
Copy pathtrain.py
File metadata and controls
72 lines (59 loc) · 3.3 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
from model import Captor
import torch
from torchvision import transforms
from dataset import MyCoco, load_vocab, collate_fn
import argparse
import random
import os
def main(args):
use_cuda = torch.cuda.is_available()
torch.manual_seed(random.randint(1, 10000))
device = torch.device("cuda" if use_cuda else "cpu")
kwargs = {'collate_fn': collate_fn, 'num_workers': 1, 'pin_memory': True} if use_cuda else {}
words = load_vocab()
train_loader = torch.utils.data.DataLoader(
MyCoco(words, args.root_dir, args.anno_path,
transform=transforms.Compose([
transforms.Resize([args.im_size] * 2),
transforms.RandomHorizontalFlip(),
transforms.ToTensor(),
transforms.Normalize(mean=[0.407, 0.457, 0.485], # subtract imagenet mean
std=[1, 1, 1]),
])),
batch_size=args.batch_size, shuffle=True, **kwargs)
test_loader = torch.utils.data.DataLoader(
MyCoco(words, args.eval_dir, args.anno_eval,
transform=transforms.Compose([
transforms.Resize([args.im_size] * 2),
transforms.ToTensor(),
transforms.Normalize(mean=[0.407, 0.457, 0.485], # subtract imagenet mean
std=[1, 1, 1])])),
batch_size=100, **kwargs)
print(len(train_loader))
model = Captor(args.lr, args.weight_decay, args.lr_decay_rate, len(words), args.embed_size)
model.to(device)
init_epoch = model.load_checkpoint(args.ckpt_path)
for i in range(1, args.epochs + 1):
model._train_ep(train_loader, device, init_epoch + i, args)
if (i + init_epoch) % args.eval_interval == 0:
model._eval_ep(test_loader, device, init_epoch + i, args)
if __name__ == '__main__':
parser = argparse.ArgumentParser(description='Image Captioning')
parser.add_argument('--im-size', nargs='?', type=int, default=224)
parser.add_argument('--batch-size', nargs='?', type=int, default=128)
parser.add_argument('--embed-size', nargs='?', type=int, default=512)
parser.add_argument('--log-interval', nargs='?', type=int, default=1)
parser.add_argument('--sv-interval', nargs='?', type=int, default=1)
parser.add_argument('--lr-decay-interval', nargs='?', type=int, default=2000)
parser.add_argument('--lr-decay-rate', nargs='?', type=float, default=5.0)
parser.add_argument('--epochs', nargs='?', type=int, default=100)
parser.add_argument('--lr', nargs='?', type=float, default=1e-3)
parser.add_argument('--weight-decay', nargs='?', type=float, default=1e-5)
parser.add_argument('--ckpt-path', nargs='?', default='checkpoints')
parser.add_argument('--root-dir', nargs='?', default=os.path.join('data', 'train2014'))
parser.add_argument('--anno-path', nargs='?', default=os.path.join('data', 'annotations', 'captions_train2014.json'))
parser.add_argument('--eval-dir', nargs='?', default=os.path.join('data', 'val2014'))
parser.add_argument('--anno-eval', nargs='?', default=os.path.join('data', 'annotations', 'captions_val2014.json'))
parser.add_argument('--eval-interval', nargs='?', type=int, default=10)
args = parser.parse_args()
main(args)