-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathrun.py
More file actions
131 lines (104 loc) · 4.79 KB
/
Copy pathrun.py
File metadata and controls
131 lines (104 loc) · 4.79 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
import torch
import numpy as np
import time
import os
import pickle
from tensorboardX import SummaryWriter
from torch.utils.data import DataLoader
from runner import Runner
from dataloader import create_dataloader, create_testloader
import torchvision.transforms as transforms
from tokenizer import Tokenizer
class Config(object):
def __init__(self) -> None:
self.run_name = "test"
self.start_epoch = 0
self.max_epoch = 100
self.lr_max = 1e-3
self.lr_min = 1e-4
self.lr_decay = pow((self.lr_min / self.lr_max), 1 / self.max_epoch)
self.device = torch.device('cuda:0')
self.lr_encoder = self.lr_max
self.lr_decoder = self.lr_max
self.train_ratio = 0.9 # valid_ratio = 1 - train_ratio
self.train_batch_size = 32
self.valid_batch_size = 1000
self.test_batch_size = 1000
self.save_interval = 5
self.log_interval = 30
self.validate_period = 5
self.batch_size = 32
self.save_period = 5
self.log_period = 30
self.max_token_length = 50
self.max_norm = 1.0
self.seed = 3407
self.img_dir = "datasets/train/images"
self.pkl_dir = "datasets/train/pkls"
self.output_dir = "output/"
# for test
self.training = True
self.test_img_dir = "datasets/test/images" # todo
self.test_ids = [i for i in range(104011, 105061)] # 0 ~ 1049
self.load_path = None # 'output/test_20241128T104341/save_model/'
self.load_model = None # 'epoch-25.pt'
self.load_name = None # 'test_20241128T104341
self.max_length = 260
def main():
config = Config()
tb_logger = None
if config.training:
config.run_name = "{}_{}".format(config.run_name, time.strftime("%Y%m%dT%H%M%S"))
config.output_dir = config.output_dir + config.run_name + '/'
if not os.path.exists(config.output_dir):
os.makedirs(config.output_dir)
tb_logger = SummaryWriter(os.path.join(f"tensorboard/{config.run_name}"), config.run_name)
else:
# only test
config.output_dir = config.output_dir + config.load_name + '/'
transform = transforms.Compose([
transforms.Resize((40, 240)), # Resize image to a fixed size
transforms.ToTensor(), # Convert image to tensor
transforms.Normalize(mean = [0.9087, 0.9083, 0.9103], std = [0.2206, 0.2212, 0.2211]) # Normalize image
])
if config.training:
dataloader, valid_dataloader, vocab = create_dataloader(img_dir = config.img_dir,
pkl_dir = config.pkl_dir,
train_ratio = config.train_ratio,
train_batch_size = config.train_batch_size,
valid_batch_size = config.valid_batch_size,
transform = transform,
seed = 42
)
tokenizer = Tokenizer(vocab)
# 保存 tokenizer
print(f'debug: max length: {dataloader.dataset.max_length}')
# print(f'debug: max length: {valid_dataloader.dataset.max_length}')
with open(config.output_dir + 'tokenizer.pkl', 'wb') as f:
pickle.dump(tokenizer, f, -1)
print(f"Successfully save tokenizer!")
test_dataloader = create_testloader(config.test_img_dir, ids = config.test_ids,
test_batch_size = config.test_batch_size,
transform = transform)
# set seed
np.random.rand(config.seed)
torch.manual_seed(config.seed)
# figure out trainer
trainer = Runner(config, tokenizer = tokenizer) # todo
if config.load_path is not None:
trainer.load(config.load_path + config.load_model)
trainer.train(dataloader, valid_dataloader, test_dataloader, tb_logger = tb_logger)
else:
# test
# load tokenizer
with open(config.output_dir + 'tokenizer.pkl', 'rb') as f:
tokenizer = pickle.load(f)
tester = Runner(config = config, tokenizer = tokenizer)
# load model
tester.load(config.load_path + config.load_model)
test_dataloader = create_testloader(config.test_img_dir, ids = config.test_ids,
test_batch_size = config.test_batch_size,
transform = transform)
tester.test(test_dataloader, config.output_dir + 'test/', 1) # 1表示 test
if __name__ == "__main__":
main()