train.py 4.4 KB

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