mirror of
https://github.com/netfun2000/dddd_trainer.git
synced 2026-08-12 23:11:36 +08:00
beta0.1
This commit is contained in:
+5
-5
@@ -120,21 +120,22 @@ class GetLoader:
|
||||
self.transform_list.append(torchvision.transforms.ToTensor())
|
||||
if self.ImageChannel == 1:
|
||||
self.transform_list.append(torchvision.transforms.Normalize(mean=[0.456],
|
||||
std=[0.224]))
|
||||
std=[0.224]))
|
||||
else:
|
||||
if self.ImageChannel != 3:
|
||||
logger.error("ImageChannel must be 1 or 3!")
|
||||
exit()
|
||||
self.transform_list.append(torchvision.transforms.Normalize(mean=[0.485, 0.456, 0.406],
|
||||
std=[0.229, 0.224, 0.225]))
|
||||
std=[0.229, 0.224, 0.225]))
|
||||
self.transform = torchvision.transforms.Compose(self.transform_list)
|
||||
tarin_loader = LoadCache(self.cache_train_path, self.path, self.word, self.ImageChannel, self.resize, self.charset)
|
||||
tarin_loader = LoadCache(self.cache_train_path, self.path, self.word, self.ImageChannel, self.resize,
|
||||
self.charset)
|
||||
val_loader = LoadCache(self.cache_val_path, self.path, self.word, self.ImageChannel, self.resize, self.charset)
|
||||
self.loaders = {
|
||||
'train': DataLoader(dataset=tarin_loader, batch_size=self.batch_size, shuffle=True, drop_last=True,
|
||||
num_workers=0, collate_fn=self.collate_to_sparse),
|
||||
'val': DataLoader(dataset=val_loader, batch_size=self.val_batch_size, shuffle=True, drop_last=True,
|
||||
num_workers=0, collate_fn=self.collate_to_sparse),
|
||||
num_workers=0, collate_fn=self.collate_to_sparse),
|
||||
}
|
||||
|
||||
def collate_to_sparse(self, batch):
|
||||
@@ -159,4 +160,3 @@ class GetLoader:
|
||||
images_pad.append(img)
|
||||
images_pad = torch.stack(images_pad, dim=0)
|
||||
return [images_pad, torch.FloatTensor(values), torch.IntTensor(shapes)]
|
||||
|
||||
|
||||
+30
-7
@@ -18,10 +18,14 @@ class Train:
|
||||
project_name)
|
||||
self.checkpoints_path = os.path.join(self.project_path, "checkpoints")
|
||||
self.models_path = os.path.join(self.project_path, "models")
|
||||
self.epoch = 0
|
||||
self.step = 0
|
||||
self.config = Config(project_name)
|
||||
self.conf = self.config.load_config()
|
||||
|
||||
self.test_step = self.conf['Train']['TEST_STEP']
|
||||
self.save_checkpoints_step = self.conf['Train']['SAVE_CHECKPOINTS_STEP']
|
||||
|
||||
self.target = self.conf['Train']['TARGET']
|
||||
self.target_acc = self.target['Accuracy']
|
||||
self.min_epoch = self.target['Epoch']
|
||||
@@ -44,14 +48,31 @@ class Train:
|
||||
self.gpu_id = -1
|
||||
self.device = self.net.get_device(self.gpu_id)
|
||||
logger.info("\nUSE CPU".format(self.gpu_id))
|
||||
|
||||
logger.info("\nSearch for history checkpoints...")
|
||||
history_checkpoints = os.listdir(self.checkpoints_path)
|
||||
if len(history_checkpoints) > 0:
|
||||
history_step = 0
|
||||
newer_checkpoint = None
|
||||
for checkpoint in history_checkpoints:
|
||||
checkpoint_name = checkpoint.split(".")[0].split("_")
|
||||
if int(checkpoint_name[2]) > history_step:
|
||||
newer_checkpoint = checkpoint
|
||||
history_step = int(checkpoint_name[2])
|
||||
self.epoch, self.step, self.lr = self.net.load_checkpoint(
|
||||
os.path.join(self.checkpoints_path, newer_checkpoint))
|
||||
self.epoch += 1
|
||||
self.step += 1
|
||||
self.net.lr = self.lr
|
||||
|
||||
else:
|
||||
logger.info("\nEmpty history checkpoints")
|
||||
logger.info("\nGet Data Loader...")
|
||||
loaders = load_cache.GetLoader(project_name)
|
||||
self.train = loaders.loaders['train']
|
||||
self.val = loaders.loaders['val']
|
||||
logger.info("\nGet Data Loader End!")
|
||||
|
||||
self.epoch = 0
|
||||
self.step = 0
|
||||
self.loss = 0
|
||||
self.avg_loss = 0
|
||||
self.start_time = time.time()
|
||||
@@ -76,14 +97,14 @@ class Train:
|
||||
str(loss), str(self.avg_loss / 100), lr
|
||||
))
|
||||
self.avg_loss = 0
|
||||
if self.step % 2000 == 0 and self.step != 0:
|
||||
if self.step % self.save_checkpoints_step == 0 and self.step != 0:
|
||||
model_path = os.path.join(self.checkpoints_path, "checkpoint_{}_{}_{}.tar".format(
|
||||
self.project_name, self.epoch, self.step,
|
||||
))
|
||||
self.net.scheduler.step()
|
||||
self.net.save_model(model_path,
|
||||
{"net": self.net.state_dict(), "optimizer": self.net.optimizer.state_dict(),
|
||||
"epoch": self.epoch, "step": self.step})
|
||||
"epoch": self.epoch, "step": self.step, "lr": lr})
|
||||
|
||||
if self.step % self.test_step == 0:
|
||||
try:
|
||||
@@ -97,7 +118,7 @@ class Train:
|
||||
test_inputs = self.net.variable_to_device(test_inputs, self.device)
|
||||
self.net = self.net.train(False)
|
||||
pred_labels, labels_list, correct_list, error_list = self.net.tester(test_inputs, test_labels,
|
||||
test_labels_length)
|
||||
test_labels_length)
|
||||
self.net = self.net.train()
|
||||
accuracy = len(correct_list) / test_inputs.shape[0]
|
||||
logger.info("{}\tEpoch: {}\tStep: {}\tLastLoss: {}\tAvgLoss: {}\tLr: {}\tAcc: {}".format(
|
||||
@@ -120,12 +141,14 @@ class Train:
|
||||
self.project_name, str(accuracy), self.epoch, self.step,
|
||||
time.localtime(self.now_time)))
|
||||
, input_names, output_names, dynamic_ax)
|
||||
logger.info("\nExport Finished!Using Time: {}min".format(str(int(int(self.now_time * 1000) - int(self.start_time * 1000)) / 60)))
|
||||
with open(os.path.join(self.models_path, "charset.json"), 'w', encoding="utf-8") as f:
|
||||
f.write(json.dumps(self.net.charset, ensure_ascii=False))
|
||||
logger.info("\nExport Finished!Using Time: {}min".format(
|
||||
str(int(int(self.now_time) - int(self.start_time)) / 60)))
|
||||
exit()
|
||||
|
||||
self.epoch += 1
|
||||
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
Train("test1")
|
||||
|
||||
Reference in New Issue
Block a user