(warmup_steps)
| 409 | |
| 410 | |
| 411 | def linear_lr_warmpup(warmup_steps): |
| 412 | assert warmup_steps > 0 |
| 413 | |
| 414 | def lr_lambda(current_step): |
| 415 | if current_step > warmup_steps: |
| 416 | return 1.0 |
| 417 | else: |
| 418 | return current_step / warmup_steps |
| 419 | |
| 420 | return lr_lambda |