(model, args)
| 5 | |
| 6 | |
| 7 | def create_optimizer(model, args): |
| 8 | from torch.nn.parallel import DistributedDataParallel as DDP |
| 9 | raw_model = model.module if isinstance(model, DDP) else model |
| 10 | |
| 11 | optimizer_name = getattr(args, "optimizer", "adamw") |
| 12 | |
| 13 | if optimizer_name == "muon": |
| 14 | return _create_muon_split(raw_model, args) |
| 15 | else: |
| 16 | return _create_adamw(raw_model, args) |
| 17 | |
| 18 | |
| 19 | def _create_adamw(model, args): |
no test coverage detected