Scale learning rate based on batch size.
(learning_rate: float, batch_size: int)
| 214 | # ============================================================================ |
| 215 | |
| 216 | def scale_lr(learning_rate: float, batch_size: int) -> float: |
| 217 | """Scale learning rate based on batch size.""" |
| 218 | return learning_rate * (batch_size * get_world_size()) / 256.0 |
| 219 | |
| 220 | |
| 221 | def setup_linear_classifiers( |
no test coverage detected