(cfg, model)
| 213 | |
| 214 | |
| 215 | def build_optimizer(cfg, model): |
| 216 | if cfg.get("filter_bias_norm_wd", False): |
| 217 | params = param_groups_weight_decay(model, weight_decay=cfg.weight_decay) |
| 218 | else: |
| 219 | params = model.parameters() |
| 220 | |
| 221 | if cfg.name == "decoupled_adamw": |
| 222 | return DecoupledAdamW(params, lr=cfg.lr, betas=list(cfg.betas), eps=cfg.eps, weight_decay=cfg.weight_decay) |
| 223 | elif cfg.name == "adamw": |
| 224 | print( |
| 225 | "INFO: You might want to increase the weight decay because in AdamW it is scaled by the lr." |
| 226 | f" Default weight decay is ``1e-2`` -> {cfg.weight_decay}. Default lr is `lr=1e-3` -> {cfg.lr}." |
| 227 | ) |
| 228 | return AdamW(params, lr=cfg.lr, betas=list(cfg.betas), eps=cfg.eps, weight_decay=cfg.weight_decay) |
| 229 | elif cfg.name == "stableadamw": |
| 230 | try: |
| 231 | if cfg.get("log_grad_norm", False): |
| 232 | from src.optimizer import StableAdamW |
| 233 | else: |
| 234 | from optimi import StableAdamW |
| 235 | except ImportError: |
| 236 | raise ImportError("Install `pip install torch-optimi` to use the StableAdamW optimizer.") |
| 237 | |
| 238 | print( |
| 239 | "INFO: You might want to increase the weight decay because in StableAdamW it is scaled by the lr." |
| 240 | f" Default weight decay is ``1e-2`` -> {cfg.weight_decay}. Default lr is `lr=1e-3` -> {cfg.lr}." |
| 241 | ) |
| 242 | return StableAdamW(params, lr=cfg.lr, betas=list(cfg.betas), eps=cfg.eps, weight_decay=cfg.weight_decay) |
| 243 | elif cfg.name == "decoupled_stableadamw": |
| 244 | try: |
| 245 | if cfg.get("log_grad_norm", False): |
| 246 | from src.optimizer import StableAdamW |
| 247 | else: |
| 248 | from optimi import StableAdamW |
| 249 | except ImportError: |
| 250 | raise ImportError("Install `pip install torch-optimi` to use the StableAdamW optimizer.") |
| 251 | |
| 252 | return StableAdamW( |
| 253 | params, |
| 254 | lr=cfg.lr, |
| 255 | betas=list(cfg.betas), |
| 256 | eps=cfg.eps, |
| 257 | weight_decay=cfg.weight_decay, |
| 258 | decouple_lr=True, |
| 259 | ) |
| 260 | else: |
| 261 | raise ValueError(f"Not sure how to build optimizer: {cfg.name}") |
| 262 | |
| 263 | |
| 264 | def get_num_tokens_in_batch_unpadded(batch: dict): |
no test coverage detected