MCPcopy Create free account
hub / github.com/chenhaoxing/DiffusionInst / build_optimizer

Method build_optimizer

train_net.py:118–164  ·  view source on GitHub ↗
(cls, cfg, model)

Source from the content-addressed store, hash-verified

116
117 @classmethod
118 def build_optimizer(cls, cfg, model):
119 params: List[Dict[str, Any]] = []
120 memo: Set[torch.nn.parameter.Parameter] = set()
121 for key, value in model.named_parameters(recurse=True):
122 if not value.requires_grad:
123 continue
124 # Avoid duplicating parameters
125 if value in memo:
126 continue
127 memo.add(value)
128 lr = cfg.SOLVER.BASE_LR
129 weight_decay = cfg.SOLVER.WEIGHT_DECAY
130 if "backbone" in key:
131 lr = lr * cfg.SOLVER.BACKBONE_MULTIPLIER
132 params += [{"params": [value], "lr": lr, "weight_decay": weight_decay}]
133
134 def maybe_add_full_model_gradient_clipping(optim): # optim: the optimizer class
135 # detectron2 doesn't have full model gradient clipping now
136 clip_norm_val = cfg.SOLVER.CLIP_GRADIENTS.CLIP_VALUE
137 enable = (
138 cfg.SOLVER.CLIP_GRADIENTS.ENABLED
139 and cfg.SOLVER.CLIP_GRADIENTS.CLIP_TYPE == "full_model"
140 and clip_norm_val > 0.0
141 )
142
143 class FullModelGradientClippingOptimizer(optim):
144 def step(self, closure=None):
145 all_params = itertools.chain(*[x["params"] for x in self.param_groups])
146 torch.nn.utils.clip_grad_norm_(all_params, clip_norm_val)
147 super().step(closure=closure)
148
149 return FullModelGradientClippingOptimizer if enable else optim
150
151 optimizer_type = cfg.SOLVER.OPTIMIZER
152 if optimizer_type == "SGD":
153 optimizer = maybe_add_full_model_gradient_clipping(torch.optim.SGD)(
154 params, cfg.SOLVER.BASE_LR, momentum=cfg.SOLVER.MOMENTUM
155 )
156 elif optimizer_type == "ADAMW":
157 optimizer = maybe_add_full_model_gradient_clipping(torch.optim.AdamW)(
158 params, cfg.SOLVER.BASE_LR
159 )
160 else:
161 raise NotImplementedError(f"no optimizer type {optimizer_type}")
162 if not cfg.SOLVER.CLIP_GRADIENTS.CLIP_TYPE == "full_model":
163 optimizer = maybe_add_gradient_clipping(cfg, optimizer)
164 return optimizer
165
166 @classmethod
167 def ema_test(cls, cfg, model, evaluators=None):

Callers 1

__init__Method · 0.95

Calls

no outgoing calls

Tested by

no test coverage detected