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

Function main

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

Source from the content-addressed store, hash-verified

14
15
16def main(args):
17 device = torch.device(args.device if torch.cuda.is_available() else "cpu")
18
19 print(args)
20 print('Start Tensorboard with "tensorboard --logdir=runs", view at 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 data_transform = {
28 "train": transforms.Compose([transforms.RandomResizedCrop(224),
29 transforms.RandomHorizontalFlip(),
30 transforms.ToTensor(),
31 transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])]),
32 "val": transforms.Compose([transforms.Resize(256),
33 transforms.CenterCrop(224),
34 transforms.ToTensor(),
35 transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])])}
36
37 # 实例化训练数据集
38 train_dataset = MyDataSet(images_path=train_images_path,
39 images_class=train_images_label,
40 transform=data_transform["train"])
41
42 # 实例化验证数据集
43 val_dataset = MyDataSet(images_path=val_images_path,
44 images_class=val_images_label,
45 transform=data_transform["val"])
46
47 batch_size = args.batch_size
48 nw = min([os.cpu_count(), batch_size if batch_size > 1 else 0, 8]) # number of workers
49 print('Using {} dataloader workers every process'.format(nw))
50 train_loader = torch.utils.data.DataLoader(train_dataset,
51 batch_size=batch_size,
52 shuffle=True,
53 pin_memory=True,
54 num_workers=nw,
55 collate_fn=train_dataset.collate_fn)
56
57 val_loader = torch.utils.data.DataLoader(val_dataset,
58 batch_size=batch_size,
59 shuffle=False,
60 pin_memory=True,
61 num_workers=nw,
62 collate_fn=val_dataset.collate_fn)
63
64 # 如果存在预训练权重则载入
65 model = create_regnet(model_name=args.model_name,
66 num_classes=args.num_classes).to(device)
67 # print(model)
68
69 if args.weights != "":
70 if os.path.exists(args.weights):
71 weights_dict = torch.load(args.weights, map_location=device)
72 load_weights_dict = {k: v for k, v in weights_dict.items()
73 if model.state_dict()[k].numel() == v.numel()}

Callers 1

train.pyFile · 0.70

Calls 5

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

Tested by

no test coverage detected