| 21 | from optimizers.sgd import SGD |
| 22 | |
| 23 | def train(epoch, model, dataloader, optimizer, args, gm = None): |
| 24 | model.train() |
| 25 | losses = [] |
| 26 | total_iters = 0 |
| 27 | start_time = time.time() |
| 28 | for idx, batch in enumerate( |
| 29 | tqdm( |
| 30 | dataloader, desc="Epoch {0}".format(epoch), disable=(not args.progress_bar) |
| 31 | ) |
| 32 | ): |
| 33 | batch = to_device(batch, args.device) |
| 34 | if args.optimizer in ["Shampoo", "kfac"]: |
| 35 | dummy_y = gm.setup_model_call(model, batch["source"]) |
| 36 | gm.setup_loss_call(model.loss, dummy_y, batch["target"], batch["mask"]) |
| 37 | outputs, loss = gm.forward_and_backward() |
| 38 | torch.nn.utils.clip_grad_norm_(model.parameters(), |
| 39 | args.clip_norm) |
| 40 | else: |
| 41 | optimizer.zero_grad() |
| 42 | log_probas = model(batch["source"]) |
| 43 | loss = model.loss(log_probas, batch["target"], batch["mask"]) |
| 44 | losses.append(loss.item() * batch["mask"].sum().item()) |
| 45 | if args.optimizer == 'AdaHessian': |
| 46 | loss.backward(create_graph=True) |
| 47 | else: |
| 48 | loss.backward() |
| 49 | optimizer.step() |
| 50 | total_iters += 1 |
| 51 | if idx % args.print_every == 0: |
| 52 | tqdm.write(f"[TRAIN] Epoch: {epoch}, Iter : {idx} / {len(dataloader)}, Loss: {loss.item():.5f}") |
| 53 | |
| 54 | mean_loss = np.mean(losses) |
| 55 | mean_loss /= args.batch_size * dataloader.dataset.max_length |
| 56 | perplexity = math.exp(mean_loss) |
| 57 | tqdm.write(f"== [TRAIN] Epoch: {epoch}, Perplexity: {perplexity:.3f} ==>") |
| 58 | return mean_loss, perplexity, time.time() - start_time |
| 59 | |
| 60 | |
| 61 | def evaluate(epoch, model, dataloader, args, mode="val"): |