Team Ai
Apppublic

sundea/text-classification

sourceHugging Faceupdated 3y agoView on Hugging Face
0likes
train_eval.py119 linesDownload Raw Back to root
1# coding: UTF-82import numpy as np3import torch4import torch.nn as nn5import torch.nn.functional as F6from sklearn import metrics7import time8from utils import get_time_dif9from tensorboardX import SummaryWriter10 11 12# 权重初始化,默认xavier13def init_network(model, method='xavier', exclude='embedding', seed=123):14    for name, w in model.named_parameters():15        if exclude not in name:16            if 'weight' in name:17                if method == 'xavier':18                    nn.init.xavier_normal_(w)19                elif method == 'kaiming':20                    nn.init.kaiming_normal_(w)21                else:22                    nn.init.normal_(w)23            elif 'bias' in name:24                nn.init.constant_(w, 0)25            else:26                pass27 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  # 记录进行到多少batch37    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_loss58                    torch.save(model.state_dict(), config.save_path)59                    improve = '*'60                    last_improve = total_batch61                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 += 172            if total_batch - last_improve > config.require_improvement:73                # 验证集loss超过1000batch没下降,结束训练74                print("No optimization for a long time, auto-stopping...")75                flag = True76                break77        if flag:78            break79    writer.close()80    test(config, model, test_iter)81 82 83def test(config, model, test_iter):84    # test85    model.load_state_dict(torch.load(config.save_path))86    model.eval()87    start_time = time.time()88    test_acc, test_loss, test_report, test_confusion = evaluate(config, model, test_iter, test=True)89    msg = 'Test Loss: {0:>5.2},  Test Acc: {1:>6.2%}'90    print(msg.format(test_loss, test_acc))91    print("Precision, Recall and F1-Score...")92    print(test_report)93    print("Confusion Matrix...")94    print(test_confusion)95    time_dif = get_time_dif(start_time)96    print("Time usage:", time_dif)97 98 99def evaluate(config, model, data_iter, test=False):100    model.eval()101    loss_total = 0102    predict_all = np.array([], dtype=int)103    labels_all = np.array([], dtype=int)104    with torch.no_grad():105        for texts, labels in data_iter:106            outputs = model(texts)107            loss = F.cross_entropy(outputs, labels)108            loss_total += loss109            labels = labels.data.cpu().numpy()110            predic = torch.max(outputs.data, 1)[1].cpu().numpy()111            labels_all = np.append(labels_all, labels)112            predict_all = np.append(predict_all, predic)113 114    acc = metrics.accuracy_score(labels_all, predict_all)115    if test:116        report = metrics.classification_report(labels_all, predict_all, target_names=config.class_list, digits=4)117        confusion = metrics.confusion_matrix(labels_all, predict_all)118        return acc, loss_total / len(data_iter), report, confusion119    return acc, loss_total / len(data_iter)