train.py 2.8 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172
  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 lenet import LeNet5
  7. import tqdm
  8. # 指定日志目录
  9. writer = SummaryWriter(log_dir='D:/gitcode/cnn_dropout/logs')
  10. # 1. 判断是否使用CUDA
  11. device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
  12. print(f'Using device: {device}')
  13. # 2. 准备数据
  14. train_set = datasets.CIFAR10(root='D:/gitcode/cnn_dropout/data', train=True, download=True, transform=transforms.ToTensor())
  15. test_set = datasets.CIFAR10(root='D:/gitcode/cnn_dropout/data', train=False, download=True, transform=transforms.ToTensor())
  16. train_loader = DataLoader(dataset=train_set, batch_size=100, shuffle=True)
  17. test_loader = DataLoader(dataset=test_set, batch_size=100, shuffle=False)
  18. # 3. 创建模型
  19. model = LeNet5()
  20. model = model.to(device)
  21. # 4. 确定损失函数
  22. loss_fn = nn.CrossEntropyLoss()
  23. # 5. 创建优化器 使用梯度下降算法 参数的更新
  24. opt = torch.optim.Adam(model.parameters())
  25. for epoch in range(100):
  26. # 6. 训练模型
  27. model.train()
  28. train_total_loss = 0
  29. for images, labels in tqdm.tqdm(train_loader, desc="train", total=len(train_loader)):
  30. # 将数据移动到设备
  31. images, labels = images.to(device), labels.to(device)
  32. outputs = model(images) # 前向传播
  33. loss = loss_fn(outputs, labels) # 计算损失
  34. opt.zero_grad() # 清空梯度
  35. loss.backward() # 反向传播 计算梯度
  36. opt.step() # 更新参数
  37. train_total_loss += loss.item()
  38. train_avg_loss = train_total_loss / len(train_loader)
  39. print(f'Epoch {epoch+1}, Loss: {train_avg_loss:.4f}')
  40. # 7. 测试模型
  41. model.eval()
  42. test_total_loss = 0
  43. test_total_acc = 0
  44. with torch.inference_mode():
  45. for images, labels in tqdm.tqdm(test_loader, desc="test", total=len(test_loader)):
  46. images, labels = images.to(device), labels.to(device)
  47. outputs = model(images)
  48. # 计算损失
  49. loss = loss_fn(outputs, labels)
  50. test_total_loss += loss.item()
  51. pred = torch.argmax(outputs, dim=1)
  52. acc = torch.eq(pred, labels).float().mean()
  53. test_total_acc += acc.item()
  54. test_avg_acc = test_total_acc / len(test_loader)
  55. print(f'Epoch {epoch+1} Test Accuracy: {test_avg_acc:.4f}')
  56. test_avg_loss = test_total_loss / len(test_loader)
  57. print(f'Epoch {epoch+1} Test Loss: {test_avg_loss:.4f}')
  58. # 记录数据
  59. writer.add_scalars('Loss/train', {'train_avg_loss': train_avg_loss,
  60. 'test_avg_loss': test_avg_loss}, epoch)
  61. writer.add_scalar('Accuracy/test', test_avg_acc, epoch)
  62. # 8. 保存模型
  63. torch.save(model.state_dict(), 'D:/gitcode/cnn_dropout/model/lenet5.pth')