train.py 4.5 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103
  1. import torch
  2. import torch.nn as nn
  3. from torch.utils.data import DataLoader
  4. from torch.utils.tensorboard import SummaryWriter
  5. from torchvision import datasets, transforms
  6. from net import FullyConnectedNet
  7. import torch.nn.functional as F
  8. import tqdm #让循环在运行时,自动在控制台显示一个动态更新的进度条
  9. import os
  10. # 指定日志目录
  11. writer = SummaryWriter(log_dir='./mlp_mse/logs')
  12. # 1. 判断是否使用CUDA
  13. device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
  14. print(f'Using device: {device}')
  15. # 2. 准备数据,Dataset 读数据,数据预处理
  16. train_set = datasets.MNIST(root='./mlp_mse/data',
  17. train=True, download=True, transform=transforms.ToTensor())
  18. test_set = datasets.MNIST(root='./mlp_mse/data',
  19. train=False, download=True, transform=transforms.ToTensor())
  20. #DataLoader 负责打包、分批、迭代、打乱数据
  21. #batch_size 每个批次多少个样本,根据显存大小设置
  22. train_loader = DataLoader(dataset=train_set, batch_size=100, shuffle=True)
  23. test_loader = DataLoader(dataset=test_set, batch_size=100, shuffle=False)
  24. # 3. 创建模型
  25. model = FullyConnectedNet()
  26. model = model.to(device)
  27. # 4. 确定损失函数
  28. loss_fn = nn.MSELoss()
  29. # 5. 创建优化器 使用梯度下降算法 参数的更新
  30. opt = torch.optim.Adam(model.parameters())
  31. # 早停法参数:当测试集损失连续 patience 轮不再变好时,提前终止训练
  32. patience = 5
  33. best_test_loss = float('inf') # 记录到目前为止最好的测试损失
  34. no_improve_count = 0 # 记录测试损失连续没有变好的轮数
  35. max_epochs = 1000
  36. os.makedirs('./mlp/model', exist_ok=True) # 确保保存模型的目录存在
  37. for epoch in range(max_epochs):
  38. # 6. 训练模型
  39. model.train()
  40. train_total_loss = 0
  41. for images, labels in tqdm.tqdm(train_loader, desc="train", total=len(train_loader)):
  42. # 将数据移动到设备
  43. images, labels = images.to(device), labels.to(device)
  44. labels = F.one_hot(labels, num_classes=10).float()
  45. outputs = model(images) # 前向传播
  46. loss = loss_fn(outputs, labels) # 计算损失
  47. opt.zero_grad() # 清空梯度
  48. loss.backward() # 反向传播 计算梯度
  49. opt.step() # 更新参数
  50. train_total_loss += loss.item()
  51. train_avg_loss = train_total_loss / len(train_loader)
  52. print(f'Epoch {epoch+1}/{max_epochs}, Loss: {train_avg_loss:.4f}')
  53. # 7. 测试模型
  54. model.eval()
  55. test_total_loss = 0 #一轮测试的总损失
  56. test_total_acc = 0 #一轮测试的总得分
  57. #禁用梯度计算
  58. #在推理时,我们不需要反向传播,因此不需要计算损失函数对参数的梯度,节省内存、加快推理速度
  59. with torch.inference_mode():
  60. for images, labels in tqdm.tqdm(test_loader, desc="test", total=len(test_loader)):
  61. images, labels = images.to(device), labels.to(device)
  62. labels = F.one_hot(labels, num_classes=10).float()
  63. outputs = model(images)
  64. # 计算损失
  65. loss = loss_fn(outputs, labels)
  66. test_total_loss += loss.item()
  67. pred = torch.argmax(outputs, dim=1)
  68. target = torch.argmax(labels, dim=1)
  69. acc = torch.eq(pred, target).float().mean()
  70. test_total_acc += acc.item()
  71. test_avg_acc = test_total_acc / len(test_loader)
  72. print(f'epoch:{epoch+1}, Test Accuracy: {test_avg_acc:.4f}')
  73. test_avg_loss = test_total_loss / len(test_loader)
  74. print(f'epoch:{epoch+1},Test Loss: {test_avg_loss:.4f}')
  75. # 记录数据
  76. writer.add_scalars('Loss/train', {'train_avg_loss': train_avg_loss,
  77. 'test_avg_loss': test_avg_loss}, epoch)
  78. writer.add_scalar('Accuracy/test', test_avg_acc, epoch)
  79. # 8. 早停判断与模型保存
  80. if test_avg_loss < best_test_loss:
  81. # 测试损失变好了:更新最好成绩,重置计数器,保存当前最优模型
  82. best_test_loss = test_avg_loss
  83. no_improve_count = 0
  84. torch.save(model.state_dict(), './mlp/model/mnist_net_best.pth')
  85. print(f'测试损失改善,保存最优模型 (best loss: {best_test_loss:.4f})')
  86. else:
  87. # 测试损失没有变好:计数器加一
  88. no_improve_count += 1
  89. print(f'测试损失未改善 ({no_improve_count}/{patience})')
  90. if no_improve_count >= patience:
  91. # 连续 patience 轮没有变好,提前终止训练
  92. print(f'早停:测试损失连续 {patience} 轮未改善,在第 {epoch+1} 轮终止训练')
  93. break