train.py 3.7 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596
  1. from torch.utils.data import DataLoader
  2. import torch
  3. from torch import nn
  4. from torch.utils.tensorboard import SummaryWriter
  5. from data import NewsDataset
  6. from utils import build_vocab,collate_batch
  7. from net import TextClassifier
  8. writer = SummaryWriter("./text_cls/logs")
  9. DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu")
  10. #NLP任务的训练和测试流程,与计算机视觉任务相差不大
  11. train_dataset = NewsDataset(True)
  12. test_dataset = NewsDataset(False)
  13. #词表构建
  14. text_vocab=build_vocab(train_dataset)
  15. #打印词表长度
  16. print("text vocab:",len(text_vocab))
  17. #回调函数,用于不同长度的文本进行填充
  18. collate = lambda batch: collate_batch(batch,text_vocab)
  19. #小批量读取数据
  20. train_dataloader=DataLoader(train_dataset,batch_size=22000,shuffle=True, collate_fn=collate)
  21. test_dataloader = DataLoader(test_dataset,batch_size=7600,shuffle=False, collate_fn=collate)
  22. #定义模型参数
  23. vocab_size = len(text_vocab)
  24. embed_dim = 256
  25. num_classes =4
  26. padding_idx = text_vocab['<pad>']
  27. #定义模型
  28. model = TextClassifier(vocab_size,
  29. embed_dim,
  30. num_classes,
  31. padding_idx)
  32. model=model.to(DEVICE)
  33. opt = torch.optim.Adam(model.parameters())
  34. loss_fn = nn.CrossEntropyLoss()
  35. print("begin train:")
  36. # 早停法参数:当测试集损失连续 patience 轮不再变好时,提前终止训练
  37. patience = 6
  38. best_test_loss = float('inf') # 记录到目前为止最好的测试损失
  39. no_improve_count = 0 # 记录测试损失连续没有变好的轮数
  40. for epoch in range(2000):
  41. #训练
  42. train_sum_loss=0
  43. model.train()
  44. for i,(text,label) in enumerate(train_dataloader):
  45. text = text.to(DEVICE)
  46. label = label.to(DEVICE)
  47. out = model(text)
  48. loss =loss_fn(out,label)
  49. opt.zero_grad()
  50. loss.backward()
  51. opt.step()
  52. train_sum_loss+=loss.item()
  53. train_avg_loss = train_sum_loss/len(train_dataloader)
  54. #测试
  55. test_sum_acc =0
  56. test_total_loss = 0
  57. model.eval()
  58. with torch.inference_mode():
  59. for i,(text,label) in enumerate(test_dataloader):
  60. text = text.to(DEVICE)
  61. label = label.to(DEVICE)
  62. out = model(text)
  63. # 计算损失
  64. loss = loss_fn(out, label)
  65. test_total_loss += loss.item()
  66. preds = torch.argmax(out, dim=1)
  67. acc = torch.mean(torch.eq(preds,label).to(torch.float32))
  68. test_sum_acc+=acc.item()
  69. test_avg_acc = test_sum_acc / len(test_dataloader)
  70. test_avg_loss = test_total_loss / len(test_dataloader)
  71. print(f"epoch:{epoch+1} train_avg_loss:{train_avg_loss} test_avg_loss:{test_avg_loss} test_avg_acc:{test_avg_acc}")
  72. writer.add_scalars("loss", {"train": train_avg_loss, "test": test_avg_loss}, epoch)
  73. writer.add_scalar("test_avg_acc", test_avg_acc, epoch)
  74. #早停判断与模型保存
  75. if test_avg_loss < best_test_loss:
  76. # 测试损失变好了:更新最好成绩,重置计数器,保存当前最优模型
  77. best_test_loss = test_avg_loss
  78. no_improve_count = 0
  79. torch.save(model.state_dict(), './text_cls/model/net_best.pth')
  80. print(f'测试损失改善,保存最优模型 (best loss: {best_test_loss:.4f})')
  81. else:
  82. # 测试损失没有变好:计数器加一
  83. no_improve_count += 1
  84. print(f'测试损失未改善 ({no_improve_count}/{patience})')
  85. if no_improve_count >= patience:
  86. # 连续 patience 轮没有变好,提前终止训练
  87. print(f'早停:测试损失连续 {patience} 轮未改善,在第 {epoch+1} 轮终止训练')
  88. break