MCPcopy Create free account
hub / github.com/Sirwenhao/Deep-Learning-Notes / main

Function main

CV/Pytorch_classification/EfficientNet/train.py:16–105  ·  view source on GitHub ↗
(args)

Source from the content-addressed store, hash-verified

14from utils import read_split_data, train_one_epoch, evaluate
15
16def main(args):
17 device = torch.device(args.device if torch.cuda.is_available() else "cpu")
18 print(args)
19
20 print('Start Tensotbaord with "tensorboard --logdir=runs", view ar http://localhost:6006/')
21 tb_writer = SummaryWriter()
22 if os.path.exists("./weights") is False:
23 os.makedirs("./weights")
24
25 train_images_path, train_images_label, val_images_path, val_images_label = read_split_data(args.data_path)
26
27 img_size = {"B0": 224,
28 "B1": 240,
29 "B2": 260,
30 "B3": 300,
31 "B4": 380,
32 "B5": 456,
33 "B6": 528,
34 "B7": 600}
35 num_model = "B0"
36
37 data_transform = {
38 "train": transforms.Compose([transforms.RandomResizedCrop(img_size[num_model]),
39 transforms.RandomHorizontalFlip(),
40 transforms.ToTensor(),
41 transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])]),
42 "val": transforms.Compose([transforms.Resize(img_size[num_model]),
43 transforms.CenterCrop(img_size[num_model]),
44 transforms.ToTensor(),
45 transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])])}
46
47 # 实例化训练数据集
48 train_dataset = MyDataSet(imgaes_path=train_images_path,
49 images_class=train_images_label,
50 transform=data_transform["train"])
51
52 # 实例化验证数据集
53 val_dataset = MyDataSet(images_path=val_images_path,
54 images_class=val_images_label,
55 transform=data_transform["val"])
56
57 batch_size = args.batch_size
58 nw = min([os.cpu_count(), batch_size if batch_size > 1 else 0, 8])
59 print('Using {} Dataloader workers every process'.format(nw))
60 train_loader = torch.utils.data.DataLoader(train_dataset,
61 batch_size=batch_size,
62 shuffle=True,
63 num_workers=nw,
64 collate_fn=train_dataset.collate_fn)
65
66 val_loader = torch.utils.data.DataLoader(val_dataset,
67 batch_size=batch_size,
68 shuffle=False,
69 pin_memory=True,
70 num_workers=nw,
71 collate_fn=val_dataset.collate_fn)
72
73 # 实例化模型加载权重

Callers 1

train.pyFile · 0.70

Calls 4

read_split_dataFunction · 0.90
MyDataSetClass · 0.90
train_one_epochFunction · 0.90
evaluateFunction · 0.90

Tested by

no test coverage detected