(args)
| 14 | |
| 15 | |
| 16 | def 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)) |
no test coverage detected