对 weight_decay 进行余弦“增加”调度,从 initial_wd -> final_wd。 参数: - optimizer: 任何带有 'weight_decay' param_group 的 optimizer - max_iters: 总的调度步数 - initial_wd: 第 0 步时的 weight_decay - final_wd: 第 max_iters 步时的 weight_decay - last_epoch: 如需从中途恢复训练,传
(
self,
optimizer,
max_iters: int,
initial_wd: float = 0.05,
final_wd: float = 0.20,
last_epoch: int = -1,
)
| 194 | |
| 195 | class CosineWeightDecayScheduler(LRScheduler): |
| 196 | def __init__( |
| 197 | self, |
| 198 | optimizer, |
| 199 | max_iters: int, |
| 200 | initial_wd: float = 0.05, |
| 201 | final_wd: float = 0.20, |
| 202 | last_epoch: int = -1, |
| 203 | ): |
| 204 | """ |
| 205 | 对 weight_decay 进行余弦“增加”调度,从 initial_wd -> final_wd。 |
| 206 | |
| 207 | 参数: |
| 208 | - optimizer: 任何带有 'weight_decay' param_group 的 optimizer |
| 209 | - max_iters: 总的调度步数 |
| 210 | - initial_wd: 第 0 步时的 weight_decay |
| 211 | - final_wd: 第 max_iters 步时的 weight_decay |
| 212 | - last_epoch: 如需从中途恢复训练,传入上次迭代 idx |
| 213 | """ |
| 214 | self.max_iters = max_iters |
| 215 | self.initial_wd = initial_wd |
| 216 | self.final_wd = final_wd |
| 217 | super().__init__(optimizer, last_epoch) |
| 218 | |
| 219 | def get_lr(self): |
| 220 | step = self._step_count |