MCPcopy Create free account
hub / github.com/JieCaoSec/FastTraffic / train

Function train

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

Source from the content-addressed store, hash-verified

30
31
32def train(config, model, train_iter, dev_iter, test_iter):
33
34 start_time = time.time()
35 model.train()
36 optimizer = torch.optim.Adam(model.parameters(), lr=config.learning_rate)
37
38 # scheduler = torch.optim.lr_scheduler.ExponentialLR(optimizer, gamma=0.9)
39 total_batch = 0
40 dev_best_loss = float('inf')
41 last_improve = 0
42 flag = False
43
44 for epoch in range(config.num_epochs):
45 print('Epoch [{}/{}]'.format(epoch + 1, config.num_epochs))
46 # scheduler.step()
47 for i, (trains, labels) in enumerate(train_iter):
48 s = time.time()
49 outputs = model(trains)
50 model.zero_grad()
51 loss = F.cross_entropy(outputs, labels)
52 loss.backward()
53 optimizer.step()
54 e = time.time()
55 #print(e-s)
56 if total_batch % 200 == 0:
57 true = labels.data.cpu()
58 predic = torch.max(outputs.data, 1)[1].cpu()
59 train_acc = metrics.accuracy_score(true, predic)
60 dev_acc, dev_loss = evaluate(config, model, dev_iter)
61 if dev_loss < dev_best_loss:
62 dev_best_loss = dev_loss
63 torch.save(model.state_dict(), config.save_path)
64 improve = '*'
65 last_improve = total_batch
66 else:
67 improve = ''
68 time_dif = get_time_dif(start_time)
69 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}'
70 print(msg.format(total_batch, loss.item(), train_acc, dev_loss, dev_acc, time_dif, improve))
71 model.train()
72
73
74 total_batch += 1
75 if total_batch - last_improve > config.require_improvement:
76 print("No optimization for a long time, auto-stopping...")
77 flag = True
78 break
79 if flag:
80 break
81
82 test(config, model, test_iter)
83
84
85

Callers 1

mainFunction · 0.90

Calls 3

get_time_difFunction · 0.90
evaluateFunction · 0.85
testFunction · 0.85

Tested by

no test coverage detected