(model, optimizer, dataloader, steps, log_interval)
| 112 | |
| 113 | |
| 114 | def train_loop(model, optimizer, dataloader, steps, log_interval): |
| 115 | model.train() |
| 116 | history = [] |
| 117 | step = 0 |
| 118 | t0 = time.time() |
| 119 | |
| 120 | while step < steps: |
| 121 | for batch in dataloader: |
| 122 | if step >= steps: |
| 123 | break |
| 124 | |
| 125 | input_ids = batch["input_ids"] |
| 126 | attention_mask = batch["attention_mask"] |
| 127 | labels = batch["labels"] |
| 128 | |
| 129 | outputs = model(input_ids=input_ids, attention_mask=attention_mask, labels=labels) |
| 130 | loss = outputs.loss |
| 131 | loss.backward() |
| 132 | |
| 133 | optimizer.step() |
| 134 | optimizer.zero_grad() |
| 135 | |
| 136 | loss_val = loss.item() |
| 137 | elapsed = time.time() - t0 |
| 138 | history.append((step, loss_val, elapsed)) |
| 139 | |
| 140 | if step % log_interval == 0: |
| 141 | print(f" step {step:4d} | loss {loss_val:.4f} | time {elapsed:.1f}s") |
| 142 | |
| 143 | step += 1 |
| 144 | |
| 145 | return history |
| 146 | |
| 147 | |
| 148 | def get_torch_dtype(name): |
no test coverage detected