(config, model, train_iter, dev_iter, test_iter)
| 27 | |
| 28 | |
| 29 | def train(config, model, train_iter, dev_iter, test_iter): |
| 30 | start_time = time.time() |
| 31 | model.train() |
| 32 | optimizer = torch.optim.Adam(model.parameters(), lr=config.learning_rate) |
| 33 | |
| 34 | # 学习率指数衰减,每次epoch:学习率 = gamma * 学习率 |
| 35 | # scheduler = torch.optim.lr_scheduler.ExponentialLR(optimizer, gamma=0.9) |
| 36 | total_batch = 0 # 记录进行到多少batch |
| 37 | dev_best_loss = float('inf') |
| 38 | last_improve = 0 # 记录上次验证集loss下降的batch数 |
| 39 | flag = False # 记录是否很久没有效果提升 |
| 40 | writer = SummaryWriter(log_dir=config.log_path + '/' + time.strftime('%m-%d_%H.%M', time.localtime())) |
| 41 | for epoch in range(config.num_epochs): |
| 42 | print('Epoch [{}/{}]'.format(epoch + 1, config.num_epochs)) |
| 43 | # scheduler.step() # 学习率衰减 |
| 44 | for i, (trains, labels) in enumerate(train_iter): |
| 45 | outputs = model(trains) |
| 46 | model.zero_grad() |
| 47 | loss = F.cross_entropy(outputs, labels) |
| 48 | loss.backward() |
| 49 | optimizer.step() |
| 50 | if total_batch % 100 == 0: |
| 51 | # 每多少轮输出在训练集和验证集上的效果 |
| 52 | true = labels.data.cpu() |
| 53 | predic = torch.max(outputs.data, 1)[1].cpu() |
| 54 | train_acc = metrics.accuracy_score(true, predic) |
| 55 | dev_acc, dev_loss = evaluate(config, model, dev_iter) |
| 56 | if dev_loss < dev_best_loss: |
| 57 | dev_best_loss = dev_loss |
| 58 | torch.save(model.state_dict(), config.save_path) |
| 59 | improve = '*' |
| 60 | last_improve = total_batch |
| 61 | else: |
| 62 | improve = '' |
| 63 | time_dif = get_time_dif(start_time) |
| 64 | msg = 'Iter: {0:>6}, Train Loss: {1:>5.2}, Train Acc: {2:>6.2%}, Val Loss: {3:>5.2}, Val Acc: {4:>6.2%}, Time: {5} {6}' |
| 65 | print(msg.format(total_batch, loss.item(), train_acc, dev_loss, dev_acc, time_dif, improve)) |
| 66 | writer.add_scalar("loss/train", loss.item(), total_batch) |
| 67 | writer.add_scalar("loss/dev", dev_loss, total_batch) |
| 68 | writer.add_scalar("acc/train", train_acc, total_batch) |
| 69 | writer.add_scalar("acc/dev", dev_acc, total_batch) |
| 70 | model.train() |
| 71 | total_batch += 1 |
| 72 | if total_batch - last_improve > config.require_improvement: |
| 73 | # 验证集loss超过1000batch没下降,结束训练 |
| 74 | print("No optimization for a long time, auto-stopping...") |
| 75 | flag = True |
| 76 | break |
| 77 | if flag: |
| 78 | break |
| 79 | writer.close() |
| 80 | test(config, model, test_iter) |
| 81 | |
| 82 | |
| 83 | def test(config, model, test_iter): |
no test coverage detected