josedolot/HybridNet_Demo2
2
1import argparse2import datetime3import os4import traceback5 6import numpy as np7import torch8from tensorboardX import SummaryWriter9from torch import nn10from torchvision import transforms11from tqdm.autonotebook import tqdm12 13from val import val14from backbone import HybridNetsBackbone15from hybridnets.loss import FocalLoss16from utils.sync_batchnorm import patch_replication_callback17from utils.utils import replace_w_sync_bn, CustomDataParallel, get_last_weights, init_weights, boolean_string, \18 save_checkpoint, DataLoaderX, Params19from hybridnets.dataset import BddDataset20from hybridnets.loss import FocalLossSeg, TverskyLoss21from hybridnets.autoanchor import run_anchor22 23 24def get_args():25 parser = argparse.ArgumentParser('HybridNets: End-to-End Perception Network - DatVu')26 parser.add_argument('-p', '--project', type=str, default='bdd100k', help='Project file that contains parameters')27 parser.add_argument('-c', '--compound_coef', type=int, default=3, help='Coefficient of efficientnet backbone')28 parser.add_argument('-n', '--num_workers', type=int, default=12, help='Num_workers of dataloader')29 parser.add_argument('-b', '--batch_size', type=int, default=12, help='Number of images per batch among all devices')30 parser.add_argument('--freeze_backbone', type=boolean_string, default=False,31 help='Freeze encoder and neck (effnet and bifpn)')32 parser.add_argument('--freeze_det', type=boolean_string, default=False,33 help='Freeze detection head')34 parser.add_argument('--freeze_seg', type=boolean_string, default=False,35 help='Freeze segmentation head')36 parser.add_argument('--lr', type=float, default=1e-4)37 parser.add_argument('--optim', type=str, default='adamw', help='Select optimizer for training, '38 'suggest using \'admaw\' until the'39 ' very final stage then switch to \'sgd\'')40 parser.add_argument('--num_epochs', type=int, default=500)41 parser.add_argument('--val_interval', type=int, default=1, help='Number of epoches between valing phases')42 parser.add_argument('--save_interval', type=int, default=500, help='Number of steps between saving')43 parser.add_argument('--es_min_delta', type=float, default=0.0,44 help='Early stopping\'s parameter: minimum change loss to qualify as an improvement')45 parser.add_argument('--es_patience', type=int, default=0,46 help='Early stopping\'s parameter: number of epochs with no improvement after which '47 'training will be stopped. Set to 0 to disable this technique')48 parser.add_argument('--data_path', type=str, default='datasets/', help='The root folder of dataset')49 parser.add_argument('--log_path', type=str, default='checkpoints/')50 parser.add_argument('-w', '--load_weights', type=str, default=None,51 help='Whether to load weights from a checkpoint, set None to initialize,'52 'set \'last\' to load last checkpoint')53 parser.add_argument('--saved_path', type=str, default='checkpoints/')54 parser.add_argument('--debug', type=boolean_string, default=False,55 help='Whether visualize the predicted boxes of training, '56 'the output images will be in test/')57 parser.add_argument('--cal_map', type=boolean_string, default=True,58 help='Calculate mAP in validation')59 parser.add_argument('-v', '--verbose', type=boolean_string, default=True,60 help='Whether to print results per class when valing')61 parser.add_argument('--plots', type=boolean_string, default=True,62 help='Whether to plot confusion matrix when valing')63 parser.add_argument('--num_gpus', type=int, default=1,64 help='Number of GPUs to be used (0 to use CPU)')65 66 args = parser.parse_args()67 return args68 69 70class ModelWithLoss(nn.Module):71 def __init__(self, model, debug=False):72 super().__init__()73 self.criterion = FocalLoss()74 self.seg_criterion1 = TverskyLoss(mode='multilabel', alpha=0.7, beta=0.3, gamma=4.0 / 3, from_logits=False)75 self.seg_criterion2 = FocalLossSeg(mode='multilabel', alpha=0.25)76 self.model = model77 self.debug = debug78 79 def forward(self, imgs, annotations, seg_annot, obj_list=None):80 _, regression, classification, anchors, segmentation = self.model(imgs)81 82 if self.debug:83 cls_loss, reg_loss = self.criterion(classification, regression, anchors, annotations,84 imgs=imgs, obj_list=obj_list)85 tversky_loss = self.seg_criterion1(segmentation, seg_annot)86 focal_loss = self.seg_criterion2(segmentation, seg_annot)87 else:88 cls_loss, reg_loss = self.criterion(classification, regression, anchors, annotations)89 tversky_loss = self.seg_criterion1(segmentation, seg_annot)90 focal_loss = self.seg_criterion2(segmentation, seg_annot)91 92 # Visualization93 # seg_0 = seg_annot[0]94 # # print('bbb', seg_0.shape)95 # seg_0 = torch.argmax(seg_0, dim = 0)96 # # print('before', seg_0.shape)97 # seg_0 = seg_0.cpu().numpy()98 # #.transpose(1, 2, 0)99 # print(seg_0.shape)100 #101 # anh = np.zeros((384,640,3))102 #103 # anh[seg_0 == 0] = (255,0,0)104 # anh[seg_0 == 1] = (0,255,0)105 # anh[seg_0 == 2] = (0,0,255)106 #107 # anh = np.uint8(anh)108 #109 # cv2.imwrite('anh.jpg',anh)110 111 seg_loss = tversky_loss + 1 * focal_loss112 # print("TVERSKY", tversky_loss)113 # print("FOCAL", focal_loss)114 115 return cls_loss, reg_loss, seg_loss, regression, classification, anchors, segmentation116 117 118def train(opt):119 params = Params(f'projects/{opt.project}.yml')120 121 if opt.num_gpus == 0:122 os.environ['CUDA_VISIBLE_DEVICES'] = '-1'123 124 if torch.cuda.is_available():125 torch.cuda.manual_seed(42)126 else:127 torch.manual_seed(42)128 129 opt.saved_path = opt.saved_path + f'/{params.project_name}/'130 opt.log_path = opt.log_path + f'/{params.project_name}/tensorboard/'131 os.makedirs(opt.log_path, exist_ok=True)132 os.makedirs(opt.saved_path, exist_ok=True)133 134 train_dataset = BddDataset(135 params=params,136 is_train=True,137 inputsize=params.model['image_size'],138 transform=transforms.Compose([139 transforms.ToTensor(),140 transforms.Normalize(141 mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]142 )143 ])144 )145 146 training_generator = DataLoaderX(147 train_dataset,148 batch_size=opt.batch_size,149 shuffle=True,150 num_workers=opt.num_workers,151 pin_memory=params.pin_memory,152 collate_fn=BddDataset.collate_fn153 )154 155 valid_dataset = BddDataset(156 params=params,157 is_train=False,158 inputsize=params.model['image_size'],159 transform=transforms.Compose([160 transforms.ToTensor(),161 transforms.Normalize(162 mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]163 )164 ])165 )166 167 val_generator = DataLoaderX(168 valid_dataset,169 batch_size=opt.batch_size,170 shuffle=False,171 num_workers=opt.num_workers,172 pin_memory=params.pin_memory,173 collate_fn=BddDataset.collate_fn174 )175 176 if params.need_autoanchor:177 params.anchors_scales, params.anchors_ratios = run_anchor(None, train_dataset)178 179 model = HybridNetsBackbone(num_classes=len(params.obj_list), compound_coef=opt.compound_coef,180 ratios=eval(params.anchors_ratios), scales=eval(params.anchors_scales),181 seg_classes=len(params.seg_list))182 183 # load last weights184 ckpt = {}185 # last_step = None186 if opt.load_weights:187 if opt.load_weights.endswith('.pth'):188 weights_path = opt.load_weights189 else:190 weights_path = get_last_weights(opt.saved_path)191 # try:192 # last_step = int(os.path.basename(weights_path).split('_')[-1].split('.')[0])193 # except:194 # last_step = 0195 196 try:197 ckpt = torch.load(weights_path)198 model.load_state_dict(ckpt.get('model', ckpt), strict=False)199 except RuntimeError as e:200 print(f'[Warning] Ignoring {e}')201 print(202 '[Warning] Don\'t panic if you see this, this might be because you load a pretrained weights with different number of classes. The rest of the weights should be loaded already.')203 else:204 print('[Info] initializing weights...')205 init_weights(model)206 207 print('[Info] Successfully!!!')208 209 if opt.freeze_backbone:210 def freeze_backbone(m):211 classname = m.__class__.__name__212 if classname in ['EfficientNetEncoder', 'BiFPN']: # replace backbone classname when using another backbone213 print("[Info] freezing {}".format(classname))214 for param in m.parameters():215 param.requires_grad = False216 model.apply(freeze_backbone)217 print('[Info] freezed backbone')218 219 if opt.freeze_det:220 def freeze_det(m):221 classname = m.__class__.__name__222 if classname in ['Regressor', 'Classifier', 'Anchors']:223 print("[Info] freezing {}".format(classname))224 for param in m.parameters():225 param.requires_grad = False226 model.apply(freeze_det)227 print('[Info] freezed detection head')228 229 if opt.freeze_seg:230 def freeze_seg(m):231 classname = m.__class__.__name__232 if classname in ['BiFPNDecoder', 'SegmentationHead']:233 print("[Info] freezing {}".format(classname))234 for param in m.parameters():235 param.requires_grad = False236 model.apply(freeze_seg)237 print('[Info] freezed segmentation head')238 239 # https://github.com/vacancy/Synchronized-BatchNorm-PyTorch240 # apply sync_bn when using multiple gpu and batch_size per gpu is lower than 4241 # useful when gpu memory is limited.242 # because when bn is disable, the training will be very unstable or slow to converge,243 # apply sync_bn can solve it,244 # by packing all mini-batch across all gpus as one batch and normalize, then send it back to all gpus.245 # but it would also slow down the training by a little bit.246 if opt.num_gpus > 1 and opt.batch_size // opt.num_gpus < 4:247 model.apply(replace_w_sync_bn)248 use_sync_bn = True249 else:250 use_sync_bn = False251 252 writer = SummaryWriter(opt.log_path + f'/{datetime.datetime.now().strftime("%Y%m%d-%H%M%S")}/')253 254 # wrap the model with loss function, to reduce the memory usage on gpu0 and speedup255 model = ModelWithLoss(model, debug=opt.debug)256 257 if opt.num_gpus > 0:258 model = model.cuda()259 if opt.num_gpus > 1:260 model = CustomDataParallel(model, opt.num_gpus)261 if use_sync_bn:262 patch_replication_callback(model)263 264 if opt.optim == 'adamw':265 optimizer = torch.optim.AdamW(model.parameters(), opt.lr)266 else:267 optimizer = torch.optim.SGD(model.parameters(), opt.lr, momentum=0.9, nesterov=True)268 # print(ckpt)269 if opt.load_weights is not None and ckpt.get('optimizer', None):270 optimizer.load_state_dict(ckpt['optimizer'])271 272 scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer, patience=3, verbose=True)273 274 epoch = 0275 best_loss = 1e5276 best_epoch = 0277 last_step = ckpt['step'] if opt.load_weights is not None and ckpt.get('step', None) else 0278 best_fitness = ckpt['best_fitness'] if opt.load_weights is not None and ckpt.get('best_fitness', None) else 0279 step = max(0, last_step)280 model.train()281 282 num_iter_per_epoch = len(training_generator)283 try:284 for epoch in range(opt.num_epochs):285 last_epoch = step // num_iter_per_epoch286 if epoch < last_epoch:287 continue288 289 epoch_loss = []290 progress_bar = tqdm(training_generator)291 for iter, data in enumerate(progress_bar):292 if iter < step - last_epoch * num_iter_per_epoch:293 progress_bar.update()294 continue295 try:296 imgs = data['img']297 annot = data['annot']298 seg_annot = data['segmentation']299 300 if opt.num_gpus == 1:301 # if only one gpu, just send it to cuda:0302 # elif multiple gpus, send it to multiple gpus in CustomDataParallel, not here303 imgs = imgs.cuda()304 annot = annot.cuda()305 seg_annot = seg_annot.cuda().long()306 307 optimizer.zero_grad()308 cls_loss, reg_loss, seg_loss, regression, classification, anchors, segmentation = model(imgs, annot,309 seg_annot,310 obj_list=params.obj_list)311 cls_loss = cls_loss.mean()312 reg_loss = reg_loss.mean()313 seg_loss = seg_loss.mean()314 315 loss = cls_loss + reg_loss + seg_loss316 if loss == 0 or not torch.isfinite(loss):317 continue318 319 loss.backward()320 # torch.nn.utils.clip_grad_norm_(model.parameters(), 0.1)321 optimizer.step()322 323 epoch_loss.append(float(loss))324 325 progress_bar.set_description(326 'Step: {}. Epoch: {}/{}. Iteration: {}/{}. Cls loss: {:.5f}. Reg loss: {:.5f}. Seg loss: {:.5f}. Total loss: {:.5f}'.format(327 step, epoch, opt.num_epochs, iter + 1, num_iter_per_epoch, cls_loss.item(),328 reg_loss.item(), seg_loss.item(), loss.item()))329 writer.add_scalars('Loss', {'train': loss}, step)330 writer.add_scalars('Regression_loss', {'train': reg_loss}, step)331 writer.add_scalars('Classfication_loss', {'train': cls_loss}, step)332 writer.add_scalars('Segmentation_loss', {'train': seg_loss}, step)333 334 # log learning_rate335 current_lr = optimizer.param_groups[0]['lr']336 writer.add_scalar('learning_rate', current_lr, step)337 338 step += 1339 340 if step % opt.save_interval == 0 and step > 0:341 save_checkpoint(model, opt.saved_path, f'hybridnets-d{opt.compound_coef}_{epoch}_{step}.pth')342 print('checkpoint...')343 344 except Exception as e:345 print('[Error]', traceback.format_exc())346 print(e)347 continue348 349 scheduler.step(np.mean(epoch_loss))350 351 if epoch % opt.val_interval == 0:352 best_fitness, best_loss, best_epoch = val(model, optimizer, val_generator, params, opt, writer, epoch,353 step, best_fitness, best_loss, best_epoch)354 except KeyboardInterrupt:355 save_checkpoint(model, opt.saved_path, f'hybridnets-d{opt.compound_coef}_{epoch}_{step}.pth')356 finally:357 writer.close()358 359 360if __name__ == '__main__':361 opt = get_args()362 train(opt)363 