Adamfan/objectdetection
0
1# YOLOR PyTorch utils2 3import datetime4import logging5import math6import os7import platform8import subprocess9import time10from contextlib import contextmanager11from copy import deepcopy12from pathlib import Path13 14import torch15import torch.backends.cudnn as cudnn16import torch.nn as nn17import torch.nn.functional as F18import torchvision19 20try:21 import thop # for FLOPS computation22except ImportError:23 thop = None24logger = logging.getLogger(__name__)25 26 27@contextmanager28def torch_distributed_zero_first(local_rank: int):29 """30 Decorator to make all processes in distributed training wait for each local_master to do something.31 """32 if local_rank not in [-1, 0]:33 torch.distributed.barrier()34 yield35 if local_rank == 0:36 torch.distributed.barrier()37 38 39def init_torch_seeds(seed=0):40 # Speed-reproducibility tradeoff https://pytorch.org/docs/stable/notes/randomness.html41 torch.manual_seed(seed)42 if seed == 0: # slower, more reproducible43 cudnn.benchmark, cudnn.deterministic = False, True44 else: # faster, less reproducible45 cudnn.benchmark, cudnn.deterministic = True, False46 47 48def date_modified(path=__file__):49 # return human-readable file modification date, i.e. '2021-3-26'50 t = datetime.datetime.fromtimestamp(Path(path).stat().st_mtime)51 return f'{t.year}-{t.month}-{t.day}'52 53 54def git_describe(path=Path(__file__).parent): # path must be a directory55 # return human-readable git description, i.e. v5.0-5-g3e25f1e https://git-scm.com/docs/git-describe56 s = f'git -C {path} describe --tags --long --always'57 try:58 return subprocess.check_output(s, shell=True, stderr=subprocess.STDOUT).decode()[:-1]59 except subprocess.CalledProcessError as e:60 return '' # not a git repository61 62 63def select_device(device='', batch_size=None):64 # device = 'cpu' or '0' or '0,1,2,3'65 s = f'YOLOR ๐ {git_describe() or date_modified()} torch {torch.__version__} ' # string66 cpu = device.lower() == 'cpu'67 if cpu:68 os.environ['CUDA_VISIBLE_DEVICES'] = '-1' # force torch.cuda.is_available() = False69 elif device: # non-cpu device requested70 os.environ['CUDA_VISIBLE_DEVICES'] = device # set environment variable71 assert torch.cuda.is_available(), f'CUDA unavailable, invalid device {device} requested' # check availability72 73 cuda = not cpu and torch.cuda.is_available()74 if cuda:75 n = torch.cuda.device_count()76 if n > 1 and batch_size: # check that batch_size is compatible with device_count77 assert batch_size % n == 0, f'batch-size {batch_size} not multiple of GPU count {n}'78 space = ' ' * len(s)79 for i, d in enumerate(device.split(',') if device else range(n)):80 p = torch.cuda.get_device_properties(i)81 s += f"{'' if i == 0 else space}CUDA:{d} ({p.name}, {p.total_memory / 1024 ** 2}MB)\n" # bytes to MB82 else:83 s += 'CPU\n'84 85 logger.info(s.encode().decode('ascii', 'ignore') if platform.system() == 'Windows' else s) # emoji-safe86 return torch.device('cuda:0' if cuda else 'cpu')87 88 89def time_synchronized():90 # pytorch-accurate time91 if torch.cuda.is_available():92 torch.cuda.synchronize()93 return time.time()94 95 96def profile(x, ops, n=100, device=None):97 # profile a pytorch module or list of modules. Example usage:98 # x = torch.randn(16, 3, 640, 640) # input99 # m1 = lambda x: x * torch.sigmoid(x)100 # m2 = nn.SiLU()101 # profile(x, [m1, m2], n=100) # profile speed over 100 iterations102 103 device = device or torch.device('cuda:0' if torch.cuda.is_available() else 'cpu')104 x = x.to(device)105 x.requires_grad = True106 print(torch.__version__, device.type, torch.cuda.get_device_properties(0) if device.type == 'cuda' else '')107 print(f"\n{'Params':>12s}{'GFLOPS':>12s}{'forward (ms)':>16s}{'backward (ms)':>16s}{'input':>24s}{'output':>24s}")108 for m in ops if isinstance(ops, list) else [ops]:109 m = m.to(device) if hasattr(m, 'to') else m # device110 m = m.half() if hasattr(m, 'half') and isinstance(x, torch.Tensor) and x.dtype is torch.float16 else m # type111 dtf, dtb, t = 0., 0., [0., 0., 0.] # dt forward, backward112 try:113 flops = thop.profile(m, inputs=(x,), verbose=False)[0] / 1E9 * 2 # GFLOPS114 except:115 flops = 0116 117 for _ in range(n):118 t[0] = time_synchronized()119 y = m(x)120 t[1] = time_synchronized()121 try:122 _ = y.sum().backward()123 t[2] = time_synchronized()124 except: # no backward method125 t[2] = float('nan')126 dtf += (t[1] - t[0]) * 1000 / n # ms per op forward127 dtb += (t[2] - t[1]) * 1000 / n # ms per op backward128 129 s_in = tuple(x.shape) if isinstance(x, torch.Tensor) else 'list'130 s_out = tuple(y.shape) if isinstance(y, torch.Tensor) else 'list'131 p = sum(list(x.numel() for x in m.parameters())) if isinstance(m, nn.Module) else 0 # parameters132 print(f'{p:12}{flops:12.4g}{dtf:16.4g}{dtb:16.4g}{str(s_in):>24s}{str(s_out):>24s}')133 134 135def is_parallel(model):136 return type(model) in (nn.parallel.DataParallel, nn.parallel.DistributedDataParallel)137 138 139def intersect_dicts(da, db, exclude=()):140 # Dictionary intersection of matching keys and shapes, omitting 'exclude' keys, using da values141 return {k: v for k, v in da.items() if k in db and not any(x in k for x in exclude) and v.shape == db[k].shape}142 143 144def initialize_weights(model):145 for m in model.modules():146 t = type(m)147 if t is nn.Conv2d:148 pass # nn.init.kaiming_normal_(m.weight, mode='fan_out', nonlinearity='relu')149 elif t is nn.BatchNorm2d:150 m.eps = 1e-3151 m.momentum = 0.03152 elif t in [nn.Hardswish, nn.LeakyReLU, nn.ReLU, nn.ReLU6]:153 m.inplace = True154 155 156def find_modules(model, mclass=nn.Conv2d):157 # Finds layer indices matching module class 'mclass'158 return [i for i, m in enumerate(model.module_list) if isinstance(m, mclass)]159 160 161def sparsity(model):162 # Return global model sparsity163 a, b = 0., 0.164 for p in model.parameters():165 a += p.numel()166 b += (p == 0).sum()167 return b / a168 169 170def prune(model, amount=0.3):171 # Prune model to requested global sparsity172 import torch.nn.utils.prune as prune173 print('Pruning model... ', end='')174 for name, m in model.named_modules():175 if isinstance(m, nn.Conv2d):176 prune.l1_unstructured(m, name='weight', amount=amount) # prune177 prune.remove(m, 'weight') # make permanent178 print(' %.3g global sparsity' % sparsity(model))179 180 181def fuse_conv_and_bn(conv, bn):182 # Fuse convolution and batchnorm layers https://tehnokv.com/posts/fusing-batchnorm-and-conv/183 fusedconv = nn.Conv2d(conv.in_channels,184 conv.out_channels,185 kernel_size=conv.kernel_size,186 stride=conv.stride,187 padding=conv.padding,188 groups=conv.groups,189 bias=True).requires_grad_(False).to(conv.weight.device)190 191 # prepare filters192 w_conv = conv.weight.clone().view(conv.out_channels, -1)193 w_bn = torch.diag(bn.weight.div(torch.sqrt(bn.eps + bn.running_var)))194 fusedconv.weight.copy_(torch.mm(w_bn, w_conv).view(fusedconv.weight.shape))195 196 # prepare spatial bias197 b_conv = torch.zeros(conv.weight.size(0), device=conv.weight.device) if conv.bias is None else conv.bias198 b_bn = bn.bias - bn.weight.mul(bn.running_mean).div(torch.sqrt(bn.running_var + bn.eps))199 fusedconv.bias.copy_(torch.mm(w_bn, b_conv.reshape(-1, 1)).reshape(-1) + b_bn)200 201 return fusedconv202 203 204def model_info(model, verbose=False, img_size=640):205 # Model information. img_size may be int or list, i.e. img_size=640 or img_size=[640, 320]206 n_p = sum(x.numel() for x in model.parameters()) # number parameters207 n_g = sum(x.numel() for x in model.parameters() if x.requires_grad) # number gradients208 if verbose:209 print('%5s %40s %9s %12s %20s %10s %10s' % ('layer', 'name', 'gradient', 'parameters', 'shape', 'mu', 'sigma'))210 for i, (name, p) in enumerate(model.named_parameters()):211 name = name.replace('module_list.', '')212 print('%5g %40s %9s %12g %20s %10.3g %10.3g' %213 (i, name, p.requires_grad, p.numel(), list(p.shape), p.mean(), p.std()))214 215 try: # FLOPS216 from thop import profile217 stride = max(int(model.stride.max()), 32) if hasattr(model, 'stride') else 32218 img = torch.zeros((1, model.yaml.get('ch', 3), stride, stride), device=next(model.parameters()).device) # input219 flops = profile(deepcopy(model), inputs=(img,), verbose=False)[0] / 1E9 * 2 # stride GFLOPS220 img_size = img_size if isinstance(img_size, list) else [img_size, img_size] # expand if int/float221 fs = ', %.1f GFLOPS' % (flops * img_size[0] / stride * img_size[1] / stride) # 640x640 GFLOPS222 except (ImportError, Exception):223 fs = ''224 225 logger.info(f"Model Summary: {len(list(model.modules()))} layers, {n_p} parameters, {n_g} gradients{fs}")226 227 228def load_classifier(name='resnet101', n=2):229 # Loads a pretrained model reshaped to n-class output230 model = torchvision.models.__dict__[name](pretrained=True)231 232 # ResNet model properties233 # input_size = [3, 224, 224]234 # input_space = 'RGB'235 # input_range = [0, 1]236 # mean = [0.485, 0.456, 0.406]237 # std = [0.229, 0.224, 0.225]238 239 # Reshape output to n classes240 filters = model.fc.weight.shape[1]241 model.fc.bias = nn.Parameter(torch.zeros(n), requires_grad=True)242 model.fc.weight = nn.Parameter(torch.zeros(n, filters), requires_grad=True)243 model.fc.out_features = n244 return model245 246 247def scale_img(img, ratio=1.0, same_shape=False, gs=32): # img(16,3,256,416)248 # scales img(bs,3,y,x) by ratio constrained to gs-multiple249 if ratio == 1.0:250 return img251 else:252 h, w = img.shape[2:]253 s = (int(h * ratio), int(w * ratio)) # new size254 img = F.interpolate(img, size=s, mode='bilinear', align_corners=False) # resize255 if not same_shape: # pad/crop img256 h, w = [math.ceil(x * ratio / gs) * gs for x in (h, w)]257 return F.pad(img, [0, w - s[1], 0, h - s[0]], value=0.447) # value = imagenet mean258 259 260def copy_attr(a, b, include=(), exclude=()):261 # Copy attributes from b to a, options to only include [...] and to exclude [...]262 for k, v in b.__dict__.items():263 if (len(include) and k not in include) or k.startswith('_') or k in exclude:264 continue265 else:266 setattr(a, k, v)267 268 269class ModelEMA:270 """ Model Exponential Moving Average from https://github.com/rwightman/pytorch-image-models271 Keep a moving average of everything in the model state_dict (parameters and buffers).272 This is intended to allow functionality like273 https://www.tensorflow.org/api_docs/python/tf/train/ExponentialMovingAverage274 A smoothed version of the weights is necessary for some training schemes to perform well.275 This class is sensitive where it is initialized in the sequence of model init,276 GPU assignment and distributed training wrappers.277 """278 279 def __init__(self, model, decay=0.9999, updates=0):280 # Create EMA281 self.ema = deepcopy(model.module if is_parallel(model) else model).eval() # FP32 EMA282 # if next(model.parameters()).device.type != 'cpu':283 # self.ema.half() # FP16 EMA284 self.updates = updates # number of EMA updates285 self.decay = lambda x: decay * (1 - math.exp(-x / 2000)) # decay exponential ramp (to help early epochs)286 for p in self.ema.parameters():287 p.requires_grad_(False)288 289 def update(self, model):290 # Update EMA parameters291 with torch.no_grad():292 self.updates += 1293 d = self.decay(self.updates)294 295 msd = model.module.state_dict() if is_parallel(model) else model.state_dict() # model state_dict296 for k, v in self.ema.state_dict().items():297 if v.dtype.is_floating_point:298 v *= d299 v += (1. - d) * msd[k].detach()300 301 def update_attr(self, model, include=(), exclude=('process_group', 'reducer')):302 # Update EMA attributes303 copy_attr(self.ema, model, include, exclude)304 305 306class BatchNormXd(torch.nn.modules.batchnorm._BatchNorm):307 def _check_input_dim(self, input):308 # The only difference between BatchNorm1d, BatchNorm2d, BatchNorm3d, etc309 # is this method that is overwritten by the sub-class310 # This original goal of this method was for tensor sanity checks311 # If you're ok bypassing those sanity checks (eg. if you trust your inference312 # to provide the right dimensional inputs), then you can just use this method313 # for easy conversion from SyncBatchNorm314 # (unfortunately, SyncBatchNorm does not store the original class - if it did315 # we could return the one that was originally created)316 return317 318def revert_sync_batchnorm(module):319 # this is very similar to the function that it is trying to revert:320 # https://github.com/pytorch/pytorch/blob/c8b3686a3e4ba63dc59e5dcfe5db3430df256833/torch/nn/modules/batchnorm.py#L679321 module_output = module322 if isinstance(module, torch.nn.modules.batchnorm.SyncBatchNorm):323 new_cls = BatchNormXd324 module_output = BatchNormXd(module.num_features,325 module.eps, module.momentum,326 module.affine,327 module.track_running_stats)328 if module.affine:329 with torch.no_grad():330 module_output.weight = module.weight331 module_output.bias = module.bias332 module_output.running_mean = module.running_mean333 module_output.running_var = module.running_var334 module_output.num_batches_tracked = module.num_batches_tracked335 if hasattr(module, "qconfig"):336 module_output.qconfig = module.qconfig337 for name, child in module.named_children():338 module_output.add_module(name, revert_sync_batchnorm(child))339 del module340 return module_output341 342 343class TracedModel(nn.Module):344 345 def __init__(self, model=None, device=None, img_size=(640,640)): 346 super(TracedModel, self).__init__()347 348 print(" Convert model to Traced-model... ") 349 self.stride = model.stride350 self.names = model.names351 self.model = model352 353 self.model = revert_sync_batchnorm(self.model)354 self.model.to('cpu')355 self.model.eval()356 357 self.detect_layer = self.model.model[-1]358 self.model.traced = True359 360 rand_example = torch.rand(1, 3, img_size, img_size)361 362 traced_script_module = torch.jit.trace(self.model, rand_example, strict=False)363 #traced_script_module = torch.jit.script(self.model)364 traced_script_module.save("traced_model.pt")365 print(" traced_script_module saved! ")366 self.model = traced_script_module367 self.model.to(device)368 self.detect_layer.to(device)369 print(" model is traced! \n") 370 371 def forward(self, x, augment=False, profile=False):372 out = self.model(x)373 out = self.detect_layer(out)374 return out