| 157 | return input_size |
| 158 | |
| 159 | def get_optimizer(self, batch_size): |
| 160 | if "optimizer" not in self.__dict__: |
| 161 | if self.warmup_epochs > 0: |
| 162 | lr = self.warmup_lr |
| 163 | else: |
| 164 | lr = self.basic_lr_per_img * batch_size |
| 165 | |
| 166 | pg0, pg1, pg2 = [], [], [] # optimizer parameter groups |
| 167 | |
| 168 | for k, v in self.model.named_modules(): |
| 169 | if hasattr(v, "bias") and isinstance(v.bias, nn.Parameter): |
| 170 | pg2.append(v.bias) # biases |
| 171 | if isinstance(v, nn.BatchNorm2d) or "bn" in k: |
| 172 | pg0.append(v.weight) # no decay |
| 173 | elif hasattr(v, "weight") and isinstance(v.weight, nn.Parameter): |
| 174 | pg1.append(v.weight) # apply decay |
| 175 | |
| 176 | optimizer = torch.optim.SGD( |
| 177 | pg0, lr=lr, momentum=self.momentum, nesterov=True |
| 178 | ) |
| 179 | optimizer.add_param_group( |
| 180 | {"params": pg1, "weight_decay": self.weight_decay} |
| 181 | ) # add pg1 with weight_decay |
| 182 | optimizer.add_param_group({"params": pg2}) |
| 183 | self.optimizer = optimizer |
| 184 | |
| 185 | return self.optimizer |
| 186 | |
| 187 | def get_lr_scheduler(self, lr, iters_per_epoch): |
| 188 | from yolox.utils import LRScheduler |