(cfg)
| 109 | |
| 110 | |
| 111 | def build_scheduler(cfg): |
| 112 | if cfg.name == "constant_with_warmup": |
| 113 | return ConstantWithWarmupScheduler(t_warmup=cfg.t_warmup) |
| 114 | elif cfg.name == "cosine_with_warmup": |
| 115 | return CosineAnnealingWithWarmupScheduler( |
| 116 | t_warmup=cfg.t_warmup, alpha_f=cfg.alpha_f |
| 117 | ) |
| 118 | elif cfg.name == "linear_decay_with_warmup": |
| 119 | return LinearWithWarmupScheduler(t_warmup=cfg.t_warmup, alpha_f=cfg.alpha_f) |
| 120 | elif cfg.name == "warmup_stable_decay": |
| 121 | return WarmupStableDecayScheduler(t_warmup=cfg.t_warmup, alpha_f=cfg.alpha_f) |
| 122 | else: |
| 123 | raise ValueError(f"Not sure how to build scheduler: {cfg.name}") |
| 124 | |
| 125 | |
| 126 | def build_model( |
no test coverage detected