| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596 |
- from torch.utils.data import DataLoader
- import torch
- from torch import nn
- from torch.utils.tensorboard import SummaryWriter
- from data import NewsDataset
- from utils import build_vocab,collate_batch
- from net import TextClassifier
- writer = SummaryWriter("./text_cls/logs")
- DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu")
- #NLP任务的训练和测试流程,与计算机视觉任务相差不大
- train_dataset = NewsDataset(True)
- test_dataset = NewsDataset(False)
- #词表构建
- text_vocab=build_vocab(train_dataset)
- #打印词表长度
- print("text vocab:",len(text_vocab))
- #回调函数,用于不同长度的文本进行填充
- collate = lambda batch: collate_batch(batch,text_vocab)
- #小批量读取数据
- train_dataloader=DataLoader(train_dataset,batch_size=22000,shuffle=True, collate_fn=collate)
- test_dataloader = DataLoader(test_dataset,batch_size=7600,shuffle=False, collate_fn=collate)
- #定义模型参数
- vocab_size = len(text_vocab)
- embed_dim = 256
- num_classes =4
- padding_idx = text_vocab['<pad>']
- #定义模型
- model = TextClassifier(vocab_size,
- embed_dim,
- num_classes,
- padding_idx)
- model=model.to(DEVICE)
- opt = torch.optim.Adam(model.parameters())
- loss_fn = nn.CrossEntropyLoss()
- print("begin train:")
- # 早停法参数:当测试集损失连续 patience 轮不再变好时,提前终止训练
- patience = 6
- best_test_loss = float('inf') # 记录到目前为止最好的测试损失
- no_improve_count = 0 # 记录测试损失连续没有变好的轮数
- for epoch in range(2000):
- #训练
- train_sum_loss=0
- model.train()
- for i,(text,label) in enumerate(train_dataloader):
- text = text.to(DEVICE)
- label = label.to(DEVICE)
-
- out = model(text)
- loss =loss_fn(out,label)
-
- opt.zero_grad()
- loss.backward()
- opt.step()
-
- train_sum_loss+=loss.item()
- train_avg_loss = train_sum_loss/len(train_dataloader)
- #测试
- test_sum_acc =0
- test_total_loss = 0
- model.eval()
- with torch.inference_mode():
- for i,(text,label) in enumerate(test_dataloader):
- text = text.to(DEVICE)
- label = label.to(DEVICE)
- out = model(text)
- # 计算损失
- loss = loss_fn(out, label)
- test_total_loss += loss.item()
-
- preds = torch.argmax(out, dim=1)
- acc = torch.mean(torch.eq(preds,label).to(torch.float32))
- test_sum_acc+=acc.item()
-
- test_avg_acc = test_sum_acc / len(test_dataloader)
- test_avg_loss = test_total_loss / len(test_dataloader)
- print(f"epoch:{epoch+1} train_avg_loss:{train_avg_loss} test_avg_loss:{test_avg_loss} test_avg_acc:{test_avg_acc}")
-
- writer.add_scalars("loss", {"train": train_avg_loss, "test": test_avg_loss}, epoch)
- writer.add_scalar("test_avg_acc", test_avg_acc, epoch)
- #早停判断与模型保存
- if test_avg_loss < best_test_loss:
- # 测试损失变好了:更新最好成绩,重置计数器,保存当前最优模型
- best_test_loss = test_avg_loss
- no_improve_count = 0
- torch.save(model.state_dict(), './text_cls/model/net_best.pth')
- print(f'测试损失改善,保存最优模型 (best loss: {best_test_loss:.4f})')
- else:
- # 测试损失没有变好:计数器加一
- no_improve_count += 1
- print(f'测试损失未改善 ({no_improve_count}/{patience})')
- if no_improve_count >= patience:
- # 连续 patience 轮没有变好,提前终止训练
- print(f'早停:测试损失连续 {patience} 轮未改善,在第 {epoch+1} 轮终止训练')
- break
|