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

Function main

CV/Pytorch_classification/ResNet/train.py:15–135  ·  view source on GitHub ↗
()

Source from the content-addressed store, hash-verified

13from model import resnet34
14
15def main():
16 device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
17 print("using {} device.".format(device))
18
19# 此处transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])中使用的是ImageNet数据集上图像RGB三通道的均值和方差
20
21 # 此处定义的是train和val阶段分别使用的数据处理方法
22 data_transform = {
23 "train": transforms.Compose([transforms.RandomResizedCrop(224),
24 transforms.RandomHorizontalFlip(),
25 transforms.ToTensor(),
26 transforms.Normalize([0.485, 0.456, 0.456], [0.229, 0.224, 0.225])]),
27 "val": transforms.Compose([transforms.Resize(256),
28 transforms.CenterCrop(224),
29 transforms.ToTensor(),
30 transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])])}
31
32
33 # data_root = os.path.abspath(os.path.join(os.getcwd(), "../..")) # 获取数据存储的根路径
34 data_root = '/Users/WH/Desktop/pytorch_classification'
35 print("current data path: {}".format(data_root))
36 image_path = os.path.join(data_root, "flower_data")
37 assert os.path.exists(image_path), "{} path does not exists.".format(image_path)
38 train_dataset = datasets.ImageFolder(root=os.path.join(image_path, "train"),
39 transform=data_transform["train"])
40 train_num = len(train_dataset)
41
42 # {'daisy':0, 'dandelion':1, 'roses':2, 'sunflower':3, 'tulips':4}
43
44 flower_list = train_dataset.class_to_idx
45 cla_dict = dict((val, key) for key, val in flower_list.items())
46 # write dict into json file
47 json_str = json.dumps(cla_dict, indent=4)
48 with open('class_indices.json', 'w') as json_file:
49 json_file.write(json_str)
50
51 batch_size = 16
52 nw = min([os.cpu_count(), batch_size if batch_size > 1 else 0, 8]) # num of workes
53 print('Using {} dataloader workers every process'.format(nw))
54
55 train_loader = torch.utils.data.DataLoader(train_dataset,
56 batch_size=batch_size, shuffle=True,
57 num_workers=nw)
58
59 validate_dataset = datasets.ImageFolder(root=os.path.join(image_path, "val"),
60 transform=data_transform["val"])
61 val_num = len(validate_dataset)
62 validate_loader = torch.utils.data.DataLoader(validate_dataset,
63 batch_size=batch_size, shuffle=False,
64 num_workers=nw)
65
66 print("using {} images for training, {} images for validation.".format(train_num,
67 val_num))
68
69
70 net = resnet34() # 模型实例化
71
72 # load pretrain weights

Callers 1

train.pyFile · 0.70

Calls 2

resnet34Function · 0.90
backwardMethod · 0.80

Tested by

no test coverage detected