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

Function main

CV/Pytorch_classification/ShuffleNet/train.py:16–109  ·  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 = shufflenet_v2_x1_0(num_classes=args.num_classes).to(device)
66 if args.weights != "Pytorch_classification/ShuffleNet/shufflenetv2_x1.pth":
67 if os.path.exists(args.weights):
68 weights_dict = torch.load(args.weights, map_location=device)
69 load_weights_dict = {k: v for k, v in weights_dict.items()
70 if model.state_dict()[k].numel() == v.numel()}
71 print(model.load_state_dict(load_weights_dict, strict=False))
72 else:
73 raise FileNotFoundError("not found weights file: {}".format(args.weights))

Callers 1

train.pyFile · 0.70

Calls 5

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

Tested by

no test coverage detected