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

Function train

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

Source from the content-addressed store, hash-verified

29
30
31def train(config, model, train_iter, dev_iter, test_iter):
32 start_time = time.time()
33 model.train()
34 param_optimizer = list(model.named_parameters())
35 no_decay = ['bias', 'LayerNorm.bias', 'LayerNorm.weight']
36 optimizer_grouped_parameters = [
37 {'params': [p for n, p in param_optimizer if not any(nd in n for nd in no_decay)], 'weight_decay': 0.01},
38 {'params': [p for n, p in param_optimizer if any(nd in n for nd in no_decay)], 'weight_decay': 0.0}]
39 # optimizer = torch.optim.Adam(model.parameters(), lr=config.learning_rate)
40 optimizer = BertAdam(optimizer_grouped_parameters,
41 lr=config.learning_rate,
42 warmup=0.05,
43 t_total=len(train_iter) * config.num_epochs)
44 total_batch = 0 # 记录进行到多少batch
45 dev_best_loss = float('inf')
46 last_improve = 0 # 记录上次验证集loss下降的batch数
47 flag = False # 记录是否很久没有效果提升
48 model.train()
49 for epoch in range(config.num_epochs):
50 print('Epoch [{}/{}]'.format(epoch + 1, config.num_epochs))
51 for i, (trains, labels) in enumerate(train_iter):
52 outputs = model(trains)
53 model.zero_grad()
54 loss = F.cross_entropy(outputs, labels)
55 loss.backward()
56 optimizer.step()
57 if total_batch % 100 == 0:
58 # 每多少轮输出在训练集和验证集上的效果
59 true = labels.data.cpu()
60 predic = torch.max(outputs.data, 1)[1].cpu()
61 train_acc = metrics.accuracy_score(true, predic)
62 dev_acc, dev_loss = evaluate(config, model, dev_iter)
63 if dev_loss < dev_best_loss:
64 dev_best_loss = dev_loss
65 torch.save(model.state_dict(), config.save_path)
66 improve = '*'
67 last_improve = total_batch
68 else:
69 improve = ''
70 time_dif = get_time_dif(start_time)
71 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}'
72 print(msg.format(total_batch, loss.item(), train_acc, dev_loss, dev_acc, time_dif, improve))
73 model.train()
74 total_batch += 1
75 if total_batch - last_improve > config.require_improvement:
76 # 验证集loss超过1000batch没下降,结束训练
77 print("No optimization for a long time, auto-stopping...")
78 flag = True
79 break
80 if flag:
81 break
82 test(config, model, test_iter)
83
84
85def test(config, model, test_iter):

Callers 1

run.pyFile · 0.90

Calls 5

stepMethod · 0.95
BertAdamClass · 0.90
get_time_difFunction · 0.90
evaluateFunction · 0.85
testFunction · 0.85

Tested by

no test coverage detected