(
learning_rate: schedule.Schedule, *, flip_sign=True
)
| 269 | |
| 270 | |
| 271 | def scale_from_learning_rate( |
| 272 | learning_rate: schedule.Schedule, *, flip_sign=True |
| 273 | ) -> schedule.ScheduleFn: |
| 274 | learning_rate_fn = schedule.as_schedule_fn(learning_rate) |
| 275 | |
| 276 | def scale_fn(step): |
| 277 | lr = learning_rate_fn(step) |
| 278 | context = current_context() |
| 279 | if context: |
| 280 | context.add_summary("lr_schedule_step", step) |
| 281 | context.add_summary("learning_rate", lr) |
| 282 | return -lr if flip_sign else lr |
| 283 | |
| 284 | return scale_fn |
| 285 | |
| 286 | |
| 287 | def per_param_scale_by_path( |
no outgoing calls
no test coverage detected