(model, name, lr)
| 89 | |
| 90 | |
| 91 | def create_optimizer(model, name, lr): |
| 92 | optim_map = { |
| 93 | "adamw": bnb.optim.AdamW, |
| 94 | "adamw8bit": bnb.optim.AdamW8bit, |
| 95 | "adamw32bit": bnb.optim.AdamW32bit, |
| 96 | "adam": bnb.optim.Adam, |
| 97 | "adam8bit": bnb.optim.Adam8bit, |
| 98 | "adam32bit": bnb.optim.Adam32bit, |
| 99 | "lion": bnb.optim.Lion, |
| 100 | "lion8bit": bnb.optim.Lion8bit, |
| 101 | "rmsprop": bnb.optim.RMSprop, |
| 102 | "rmsprop8bit": bnb.optim.RMSprop8bit, |
| 103 | "adagrad": bnb.optim.Adagrad, |
| 104 | "adagrad8bit": bnb.optim.Adagrad8bit, |
| 105 | "lamb": bnb.optim.LAMB, |
| 106 | "lars": lambda p, lr: bnb.optim.LARS(p, lr, momentum=0.9), |
| 107 | "sgd": lambda p, lr: bnb.optim.SGD(p, lr, momentum=0.9), |
| 108 | "sgd8bit": lambda p, lr: bnb.optim.SGD8bit(p, lr, momentum=0.9), |
| 109 | } |
| 110 | factory = optim_map[name] |
| 111 | return factory(model.parameters(), lr=lr) |
| 112 | |
| 113 | |
| 114 | def train_loop(model, optimizer, dataloader, steps, log_interval): |
no outgoing calls
no test coverage detected