Arulkumar03/Wheat_HEAD_Detection_Counting_ComputerVision_Model
0
1# Copyright (c) Facebook, Inc. and its affiliates.2import copy3import itertools4import logging5from collections import defaultdict6from enum import Enum7from typing import Any, Callable, Dict, Iterable, List, Optional, Set, Type, Union8import torch9from fvcore.common.param_scheduler import (10 CosineParamScheduler,11 MultiStepParamScheduler,12 StepWithFixedGammaParamScheduler,13)14 15from detectron2.config import CfgNode16from detectron2.utils.env import TORCH_VERSION17 18from .lr_scheduler import LRMultiplier, LRScheduler, WarmupParamScheduler19 20_GradientClipperInput = Union[torch.Tensor, Iterable[torch.Tensor]]21_GradientClipper = Callable[[_GradientClipperInput], None]22 23 24class GradientClipType(Enum):25 VALUE = "value"26 NORM = "norm"27 28 29def _create_gradient_clipper(cfg: CfgNode) -> _GradientClipper:30 """31 Creates gradient clipping closure to clip by value or by norm,32 according to the provided config.33 """34 cfg = copy.deepcopy(cfg)35 36 def clip_grad_norm(p: _GradientClipperInput):37 torch.nn.utils.clip_grad_norm_(p, cfg.CLIP_VALUE, cfg.NORM_TYPE)38 39 def clip_grad_value(p: _GradientClipperInput):40 torch.nn.utils.clip_grad_value_(p, cfg.CLIP_VALUE)41 42 _GRADIENT_CLIP_TYPE_TO_CLIPPER = {43 GradientClipType.VALUE: clip_grad_value,44 GradientClipType.NORM: clip_grad_norm,45 }46 return _GRADIENT_CLIP_TYPE_TO_CLIPPER[GradientClipType(cfg.CLIP_TYPE)]47 48 49def _generate_optimizer_class_with_gradient_clipping(50 optimizer: Type[torch.optim.Optimizer],51 *,52 per_param_clipper: Optional[_GradientClipper] = None,53 global_clipper: Optional[_GradientClipper] = None,54) -> Type[torch.optim.Optimizer]:55 """56 Dynamically creates a new type that inherits the type of a given instance57 and overrides the `step` method to add gradient clipping58 """59 assert (60 per_param_clipper is None or global_clipper is None61 ), "Not allowed to use both per-parameter clipping and global clipping"62 63 def optimizer_wgc_step(self, closure=None):64 if per_param_clipper is not None:65 for group in self.param_groups:66 for p in group["params"]:67 per_param_clipper(p)68 else:69 # global clipper for future use with detr70 # (https://github.com/facebookresearch/detr/pull/287)71 all_params = itertools.chain(*[g["params"] for g in self.param_groups])72 global_clipper(all_params)73 super(type(self), self).step(closure)74 75 OptimizerWithGradientClip = type(76 optimizer.__name__ + "WithGradientClip",77 (optimizer,),78 {"step": optimizer_wgc_step},79 )80 return OptimizerWithGradientClip81 82 83def maybe_add_gradient_clipping(84 cfg: CfgNode, optimizer: Type[torch.optim.Optimizer]85) -> Type[torch.optim.Optimizer]:86 """87 If gradient clipping is enabled through config options, wraps the existing88 optimizer type to become a new dynamically created class OptimizerWithGradientClip89 that inherits the given optimizer and overrides the `step` method to90 include gradient clipping.91 92 Args:93 cfg: CfgNode, configuration options94 optimizer: type. A subclass of torch.optim.Optimizer95 96 Return:97 type: either the input `optimizer` (if gradient clipping is disabled), or98 a subclass of it with gradient clipping included in the `step` method.99 """100 if not cfg.SOLVER.CLIP_GRADIENTS.ENABLED:101 return optimizer102 if isinstance(optimizer, torch.optim.Optimizer):103 optimizer_type = type(optimizer)104 else:105 assert issubclass(optimizer, torch.optim.Optimizer), optimizer106 optimizer_type = optimizer107 108 grad_clipper = _create_gradient_clipper(cfg.SOLVER.CLIP_GRADIENTS)109 OptimizerWithGradientClip = _generate_optimizer_class_with_gradient_clipping(110 optimizer_type, per_param_clipper=grad_clipper111 )112 if isinstance(optimizer, torch.optim.Optimizer):113 optimizer.__class__ = OptimizerWithGradientClip # a bit hacky, not recommended114 return optimizer115 else:116 return OptimizerWithGradientClip117 118 119def build_optimizer(cfg: CfgNode, model: torch.nn.Module) -> torch.optim.Optimizer:120 """121 Build an optimizer from config.122 """123 params = get_default_optimizer_params(124 model,125 base_lr=cfg.SOLVER.BASE_LR,126 weight_decay_norm=cfg.SOLVER.WEIGHT_DECAY_NORM,127 bias_lr_factor=cfg.SOLVER.BIAS_LR_FACTOR,128 weight_decay_bias=cfg.SOLVER.WEIGHT_DECAY_BIAS,129 )130 sgd_args = {131 "params": params,132 "lr": cfg.SOLVER.BASE_LR,133 "momentum": cfg.SOLVER.MOMENTUM,134 "nesterov": cfg.SOLVER.NESTEROV,135 "weight_decay": cfg.SOLVER.WEIGHT_DECAY,136 }137 if TORCH_VERSION >= (1, 12):138 sgd_args["foreach"] = True139 return maybe_add_gradient_clipping(cfg, torch.optim.SGD(**sgd_args))140 141 142def get_default_optimizer_params(143 model: torch.nn.Module,144 base_lr: Optional[float] = None,145 weight_decay: Optional[float] = None,146 weight_decay_norm: Optional[float] = None,147 bias_lr_factor: Optional[float] = 1.0,148 weight_decay_bias: Optional[float] = None,149 lr_factor_func: Optional[Callable] = None,150 overrides: Optional[Dict[str, Dict[str, float]]] = None,151) -> List[Dict[str, Any]]:152 """153 Get default param list for optimizer, with support for a few types of154 overrides. If no overrides needed, this is equivalent to `model.parameters()`.155 156 Args:157 base_lr: lr for every group by default. Can be omitted to use the one in optimizer.158 weight_decay: weight decay for every group by default. Can be omitted to use the one159 in optimizer.160 weight_decay_norm: override weight decay for params in normalization layers161 bias_lr_factor: multiplier of lr for bias parameters.162 weight_decay_bias: override weight decay for bias parameters.163 lr_factor_func: function to calculate lr decay rate by mapping the parameter names to164 corresponding lr decay rate. Note that setting this option requires165 also setting ``base_lr``.166 overrides: if not `None`, provides values for optimizer hyperparameters167 (LR, weight decay) for module parameters with a given name; e.g.168 ``{"embedding": {"lr": 0.01, "weight_decay": 0.1}}`` will set the LR and169 weight decay values for all module parameters named `embedding`.170 171 For common detection models, ``weight_decay_norm`` is the only option172 needed to be set. ``bias_lr_factor,weight_decay_bias`` are legacy settings173 from Detectron1 that are not found useful.174 175 Example:176 ::177 torch.optim.SGD(get_default_optimizer_params(model, weight_decay_norm=0),178 lr=0.01, weight_decay=1e-4, momentum=0.9)179 """180 if overrides is None:181 overrides = {}182 defaults = {}183 if base_lr is not None:184 defaults["lr"] = base_lr185 if weight_decay is not None:186 defaults["weight_decay"] = weight_decay187 bias_overrides = {}188 if bias_lr_factor is not None and bias_lr_factor != 1.0:189 # NOTE: unlike Detectron v1, we now by default make bias hyperparameters190 # exactly the same as regular weights.191 if base_lr is None:192 raise ValueError("bias_lr_factor requires base_lr")193 bias_overrides["lr"] = base_lr * bias_lr_factor194 if weight_decay_bias is not None:195 bias_overrides["weight_decay"] = weight_decay_bias196 if len(bias_overrides):197 if "bias" in overrides:198 raise ValueError("Conflicting overrides for 'bias'")199 overrides["bias"] = bias_overrides200 if lr_factor_func is not None:201 if base_lr is None:202 raise ValueError("lr_factor_func requires base_lr")203 norm_module_types = (204 torch.nn.BatchNorm1d,205 torch.nn.BatchNorm2d,206 torch.nn.BatchNorm3d,207 torch.nn.SyncBatchNorm,208 # NaiveSyncBatchNorm inherits from BatchNorm2d209 torch.nn.GroupNorm,210 torch.nn.InstanceNorm1d,211 torch.nn.InstanceNorm2d,212 torch.nn.InstanceNorm3d,213 torch.nn.LayerNorm,214 torch.nn.LocalResponseNorm,215 )216 params: List[Dict[str, Any]] = []217 memo: Set[torch.nn.parameter.Parameter] = set()218 for module_name, module in model.named_modules():219 for module_param_name, value in module.named_parameters(recurse=False):220 if not value.requires_grad:221 continue222 # Avoid duplicating parameters223 if value in memo:224 continue225 memo.add(value)226 227 hyperparams = copy.copy(defaults)228 if isinstance(module, norm_module_types) and weight_decay_norm is not None:229 hyperparams["weight_decay"] = weight_decay_norm230 if lr_factor_func is not None:231 hyperparams["lr"] *= lr_factor_func(f"{module_name}.{module_param_name}")232 233 hyperparams.update(overrides.get(module_param_name, {}))234 params.append({"params": [value], **hyperparams})235 return reduce_param_groups(params)236 237 238def _expand_param_groups(params: List[Dict[str, Any]]) -> List[Dict[str, Any]]:239 # Transform parameter groups into per-parameter structure.240 # Later items in `params` can overwrite parameters set in previous items.241 ret = defaultdict(dict)242 for item in params:243 assert "params" in item244 cur_params = {x: y for x, y in item.items() if x != "params"}245 for param in item["params"]:246 ret[param].update({"params": [param], **cur_params})247 return list(ret.values())248 249 250def reduce_param_groups(params: List[Dict[str, Any]]) -> List[Dict[str, Any]]:251 # Reorganize the parameter groups and merge duplicated groups.252 # The number of parameter groups needs to be as small as possible in order253 # to efficiently use the PyTorch multi-tensor optimizer. Therefore instead254 # of using a parameter_group per single parameter, we reorganize the255 # parameter groups and merge duplicated groups. This approach speeds256 # up multi-tensor optimizer significantly.257 params = _expand_param_groups(params)258 groups = defaultdict(list) # re-group all parameter groups by their hyperparams259 for item in params:260 cur_params = tuple((x, y) for x, y in item.items() if x != "params")261 groups[cur_params].extend(item["params"])262 ret = []263 for param_keys, param_values in groups.items():264 cur = {kv[0]: kv[1] for kv in param_keys}265 cur["params"] = param_values266 ret.append(cur)267 return ret268 269 270def build_lr_scheduler(cfg: CfgNode, optimizer: torch.optim.Optimizer) -> LRScheduler:271 """272 Build a LR scheduler from config.273 """274 name = cfg.SOLVER.LR_SCHEDULER_NAME275 276 if name == "WarmupMultiStepLR":277 steps = [x for x in cfg.SOLVER.STEPS if x <= cfg.SOLVER.MAX_ITER]278 if len(steps) != len(cfg.SOLVER.STEPS):279 logger = logging.getLogger(__name__)280 logger.warning(281 "SOLVER.STEPS contains values larger than SOLVER.MAX_ITER. "282 "These values will be ignored."283 )284 sched = MultiStepParamScheduler(285 values=[cfg.SOLVER.GAMMA**k for k in range(len(steps) + 1)],286 milestones=steps,287 num_updates=cfg.SOLVER.MAX_ITER,288 )289 elif name == "WarmupCosineLR":290 end_value = cfg.SOLVER.BASE_LR_END / cfg.SOLVER.BASE_LR291 assert end_value >= 0.0 and end_value <= 1.0, end_value292 sched = CosineParamScheduler(1, end_value)293 elif name == "WarmupStepWithFixedGammaLR":294 sched = StepWithFixedGammaParamScheduler(295 base_value=1.0,296 gamma=cfg.SOLVER.GAMMA,297 num_decays=cfg.SOLVER.NUM_DECAYS,298 num_updates=cfg.SOLVER.MAX_ITER,299 )300 else:301 raise ValueError("Unknown LR scheduler: {}".format(name))302 303 sched = WarmupParamScheduler(304 sched,305 cfg.SOLVER.WARMUP_FACTOR,306 min(cfg.SOLVER.WARMUP_ITERS / cfg.SOLVER.MAX_ITER, 1.0),307 cfg.SOLVER.WARMUP_METHOD,308 cfg.SOLVER.RESCALE_INTERVAL,309 )310 return LRMultiplier(optimizer, multiplier=sched, max_iter=cfg.SOLVER.MAX_ITER)311 