| 13 | from model import resnet34 |
| 14 | |
| 15 | def 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 |