Team Ai
Apppublic

Arulkumar03/Wheat_HEAD_Detection_Counting_ComputerVision_Model

sourceHugging Facemitupdated 3y agoView on Hugging Face
0likes
build.py311 linesDownload Raw Back to solver
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