|
| 1 | +import os |
| 2 | +import datetime |
| 3 | + |
| 4 | +import torch |
| 5 | + |
| 6 | +import transforms |
| 7 | +from my_dataset import VOCDataSet |
| 8 | + |
| 9 | + |
| 10 | +def create_model(num_classes=21): |
| 11 | + backbone = Backbone() # 特征提取器 |
| 12 | + model = SSD300(backbone=backbone, num_classes=num_classes) |
| 13 | + |
| 14 | + pre_ssd_path = '***' |
| 15 | + if os.path.exists(pre_ssd_path) is False: |
| 16 | + raise FileNotFoundError('*** not found in {}'.format(pre_ssd_path)) |
| 17 | + pre_model_dict = torch.load(pre_ssd_path, map_location='cpu') |
| 18 | + pre_weights_dict = pre_model_dict['model'] |
| 19 | + |
| 20 | + # 不加载类别预测器的权重,因为是voc和coco不同,但是可以使用回归预测器的权重 |
| 21 | + del_conf_loc_dict = {} |
| 22 | + for k, v in pre_weights_dict.items(): |
| 23 | + split_key = k.split('.') |
| 24 | + if 'conf' in k: |
| 25 | + continue |
| 26 | + del_conf_loc_dict.update({k: v}) |
| 27 | + |
| 28 | + missing_keys, unexpected_keys = model.load_state_dict(del_conf_loc_dict, strict=False) |
| 29 | + if len(missing_keys) != 0 or len(unexpected_keys) != 0: |
| 30 | + print("missing_keys: ", missing_keys) |
| 31 | + print("unexpected_keys: ", unexpected_keys) |
| 32 | + |
| 33 | + return model |
| 34 | + |
1 | 35 |
|
2 | 36 | def main(parser_data): |
| 37 | + # 指定GPU或CPU,Q:如何进行多GPU训练 |
| 38 | + device = torch.device(parser_data.device if torch.cuda.is_available() else 'cpu') |
| 39 | + print('Using device {} training.'.format(device.type)) |
| 40 | + |
| 41 | + if not os.path.exists('save_weights'): |
| 42 | + os.makedirs('save_weights') |
| 43 | + |
| 44 | + results_file = 'results{}.txt'.format(datetime.datetime.now().strftime('%Y%m%d-%H%M%S')) |
| 45 | + |
| 46 | + data_transform = { |
| 47 | + 'train': transforms.Compose([ |
| 48 | + transforms.SSDCropping(), # Q: 裁剪应该旨在增加训练样本,裁剪后GT也要改变? |
| 49 | + transforms.Resize(), |
| 50 | + transforms.ColorJitter(), |
| 51 | + transforms.ToTensor(), |
| 52 | + transforms.RandomHorizontalFlips(), |
| 53 | + transforms.Normalization(), |
| 54 | + # ???:输出default box集合的正样本和负样本(default_box的位置固定,正样本和GT匹配最佳且满足IOU>0.5) |
| 55 | + transforms.AssignGTtoDefaultBox() |
| 56 | + ]), |
| 57 | + 'val': transforms.Compose([ |
| 58 | + transforms.Resize(), |
| 59 | + transforms.ToTensor(), |
| 60 | + transforms.Normalization() |
| 61 | + ]) |
| 62 | + } |
| 63 | + |
| 64 | + VOC_root = parser_data.data_path |
| 65 | + if os.path.exists(os.path.join(VOC_root, 'VOCdevkit')) is False: |
| 66 | + raise FileNotFoundError('VOCdevkit does not in path: {}'.format(VOC_root)) |
| 67 | + |
| 68 | + # VOCdevkit/VOC2012/ImageSets/Main/train.txt |
| 69 | + train_dataset = VOCDataset(VOC_root, '2012', data_transform['train'], train_set='train.txt') |
| 70 | + # 训练时batch size必须大于1. |
| 71 | + batch_size = parser_data.batch_size |
| 72 | + assert batch_size > 1, 'batch size must be greater than 1' # Q:assert条件不满足,是否就推出程序? |
| 73 | + # 防止最后一个batch_size=1,如果是就舍去 |
| 74 | + drop_last = True if len(train_dataset) % batch_size == 1 else False |
| 75 | + nw = min([os.cput_count(), batch_size if batch_size > 1 else 0, 8]) # number of workers |
| 76 | + print('Using %g dataloader workers' % nw) |
| 77 | + train_data_loader = torch.utils.data.DataLoader( |
| 78 | + train_dataset, |
| 79 | + batch_size=batch_size, |
| 80 | + shuffle=True, |
| 81 | + num_workers=nw, |
| 82 | + collate_fn=train_dataset.collate_fn, # 核对函数 |
| 83 | + drop_last=drop_last |
| 84 | + ) |
| 85 | + |
| 86 | + # VOCdevkit/VOC2012/ImageSets/Main/val.txt |
| 87 | + val_dataset = VOCDataset(VOC_root, '2007', data_transform['val'], train_set='val.txt') |
| 88 | + val_data_loader = torch.utils.data.Dataloader( |
| 89 | + val_dataset, |
| 90 | + batch_size=batch_size, |
| 91 | + num_workers=nw, |
| 92 | + collate_fn=train_dataset.collate_fn # ?:为何不是val_dataset.collate_fn |
| 93 | + ) |
| 94 | + |
| 95 | + model = create_model(num_classes=parser_data.num_classes+1) |
| 96 | + model.to(device) |
| 97 | + |
| 98 | + # define optimizer |
| 99 | + params = [p for p in model.parameters() if p.requires_grad] # Q:可以更新的参数部分? |
| 100 | + optimizer = torch.optim.SGD(params, lr=0.0005, momentum=0.9, weight_decay=0.0005) |
| 101 | + |
| 102 | + # learning date scheduler,Q:step_size指每个epoch增大5倍? |
| 103 | + lr_scheduler = torch.optim.lr_scheduler.StepLR(optimizer, step_size=5, gamma=0.3) |
| 104 | + |
| 105 | + # 如果指定了上次训练保存的权重文件地址,则续接上次结果训练 |
| 106 | + if parser_data.resume != '': |
| 107 | + check_point = torch.load(parser_data.resume, map_location='cpu') # Q:在cpu上加载保存于GPU的模型? |
| 108 | + model.load_state_dict(check_point['model']) |
| 109 | + optimizer.load_state_dict(check_point['optimizer']) |
| 110 | + lr_scheduler.load_state_dict(check_point['lr_scheduler']) |
| 111 | + parser_data.start_epoch = check_point['epoch'] + 1 |
| 112 | + print('the training process from epoch {}...'.format(parser_data.start_epoch)) |
| 113 | + |
| 114 | + train_loss = [] |
| 115 | + learning_rate = [] |
| 116 | + val_map = [] |
| 117 | + |
| 118 | + # 提前加载验证集数据,以免每次验证都重新加载一次 |
| 119 | + val_data = get_coco_api_from_dataset(val_data_loader.dataset) # Q:这里是说一次性加载所有的数据? |
| 120 | + for epoch in range(parser_data.start_epoch, parser_data.epochs): |
| 121 | + mean_loss, lr = utils.train_one_epoch(model=model,optimizer=optimizer, data_loader=train_data_loader, device=device, epoch=epoch, print_freq=50) |
| 122 | + train_loss.append(mean_loss.item()) |
| 123 | + learning_rate.append(lr) |
| 124 | + |
| 125 | + # update learning rate |
| 126 | + lr_scheduler.step() |
| 127 | + |
| 128 | + # Q:是否使用data_set参数后就是一直把数据保存在GPU上,从而忽略了data_loader这种以batch加载数据的方式 |
| 129 | + coco_info = utils.evaluate(model=model, data_loader=val_data_loader, device=device, data_set=val_data) |
| 130 | + |
| 131 | + # write info txt |
| 132 | + with open(results_file, 'a') as f: |
| 133 | + result_info = [str(round(i, 4)) for i in coco_info + [mean_loss.item()]] + [str(round(lr, 6))] |
| 134 | + txt = "epoch:{} {}".format(epoch, ' '.join(result_info)) |
| 135 | + f.write(txt + "\n") |
| 136 | + |
| 137 | + val_map.append(coco_info[1]) |
| 138 | + |
| 139 | + # save weights |
| 140 | + save_files = { |
| 141 | + 'model': model.state_dict(), # Q:state_dict()是什么意思? |
| 142 | + 'optimizer': optimizer.state_dict(), |
| 143 | + 'lr_scheduler': lr_scheduler.state_dict(), |
| 144 | + 'epoch': epoch |
| 145 | + } |
| 146 | + |
| 147 | + torch.save(save_files, './save_weights/ssd300-{}.pth'.format(epoch)) |
| 148 | + |
| 149 | + # plot loss and lr curve |
| 150 | + if len(train_loss) != 0 and len(learning_rate) != 0: |
| 151 | + from plot_curve import plot_loss_and_lr |
| 152 | + plot_loss_and_lr(train_loss, lr) |
3 | 153 |
|
| 154 | + # plot mAP curve |
| 155 | + if len(val_map) != 0: |
| 156 | + from plot_curve import plot_map |
| 157 | + plot_map(val_map) |
4 | 158 |
|
5 | 159 |
|
6 | 160 | if __name__ == '__main__': |
|
0 commit comments