MCPcopy Create free account
hub / github.com/649453932/Chinese-Text-Classification-Pytorch / train

Function train

train_eval.py:29–80  ·  view source on GitHub ↗
(config, model, train_iter, dev_iter, test_iter)

Source from the content-addressed store, hash-verified

27
28
29def 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
83def test(config, model, test_iter):

Callers 1

run.pyFile · 0.90

Calls 3

get_time_difFunction · 0.90
evaluateFunction · 0.85
testFunction · 0.85

Tested by

no test coverage detected