Guest User

Untitled

a guest
Nov 7th, 2025
277
0
Never
Not a member of Pastebin yet? Sign Up, it unlocks many cool features!
text 3.35 KB | None | 0 0
  1. import torch
  2. import torch.nn as nn
  3. import torch.optim as optim
  4. import torchvision
  5. import torchvision.transforms as transforms
  6. from torch.utils.data import DataLoader
  7. from torch.optim.lr_scheduler import StepLR
  8.  
  9. import torch
  10. import torch.nn as nn
  11.  
  12. class HyperEfficientCNN(nn.Module):
  13. def __init__(self):
  14. super(HyperEfficientCNN, self).__init__()
  15. self.conv1 = nn.Conv2d(1, 8, kernel_size=3, padding=1)
  16. self.bn1 = nn.BatchNorm2d(8)
  17. self.relu1 = nn.ReLU()
  18. self.pool1 = nn.MaxPool2d(kernel_size=2, stride=2)
  19.  
  20. self.conv2 = nn.Conv2d(8, 8, kernel_size=3, padding=1)
  21. self.bn2 = nn.BatchNorm2d(8)
  22. self.relu2 = nn.ReLU()
  23. self.pool2 = nn.MaxPool2d(kernel_size=2, stride=2)
  24.  
  25. self.global_avg_pool = nn.AdaptiveAvgPool2d((1, 1))
  26. self.fc1 = nn.Linear(8, 10)
  27.  
  28. def forward(self, x):
  29. x = self.pool1(self.relu1(self.bn1(self.conv1(x))))
  30. x = self.pool2(self.relu2(self.bn2(self.conv2(x))))
  31.  
  32. x = self.global_avg_pool(x)
  33. x = x.view(x.size(0), -1)
  34. x = self.fc1(x)
  35. return x
  36.  
  37. def count_parameters(model):
  38. return sum(p.numel() for p in model.parameters() if p.requires_grad)
  39.  
  40. temp_model = HyperEfficientCNN()
  41. print(f"模型总参数量: {count_parameters(temp_model):,} 个")
  42.  
  43. device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
  44. print(f"Using device: {device}")
  45.  
  46. batch_size = 64
  47. learning_rate = 0.01
  48. num_epochs = 20
  49.  
  50. transform = transforms.Compose([
  51. transforms.ToTensor(),
  52. transforms.Normalize((0.1307,), (0.3081,))
  53. ])
  54.  
  55. train_dataset = torchvision.datasets.MNIST(root='./data', train=True, transform=transform, download=True)
  56. test_dataset = torchvision.datasets.MNIST(root='./data', train=False, transform=transform, download=True)
  57.  
  58. train_loader = DataLoader(dataset=train_dataset, batch_size=batch_size, shuffle=True)
  59. test_loader = DataLoader(dataset=test_dataset, batch_size=batch_size, shuffle=False)
  60.  
  61. model = HyperEfficientCNN().to(device)
  62.  
  63. print(f'The model has {count_parameters(model):,} trainable parameters.')
  64.  
  65. criterion = nn.CrossEntropyLoss()
  66. optimizer = optim.Adam(model.parameters(), lr=learning_rate)
  67. scheduler = StepLR(optimizer, step_size=5, gamma=0.5)
  68.  
  69. for epoch in range(num_epochs):
  70. model.train()
  71. for i, (images, labels) in enumerate(train_loader):
  72. images = images.to(device)
  73. labels = labels.to(device)
  74.  
  75. outputs = model(images)
  76. loss = criterion(outputs, labels)
  77.  
  78. optimizer.zero_grad()
  79. loss.backward()
  80. optimizer.step()
  81.  
  82. if (i + 1) % 200 == 0:
  83. print(f'Epoch [{epoch+1}/{num_epochs}], Step [{i+1}/{len(train_loader)}], Loss: {loss.item():.4f}')
  84.  
  85. scheduler.step()
  86. print(f"Epoch {epoch+1} finished. Current learning rate: {scheduler.get_last_lr()[0]}")
  87.  
  88. print("\nTraining Finished!")
  89.  
  90. model.eval()
  91. with torch.no_grad():
  92. correct = 0
  93. total = 0
  94. for images, labels in test_loader:
  95. images = images.to(device)
  96. labels = labels.to(device)
  97. outputs = model(images)
  98. _, predicted = torch.max(outputs.data, 1)
  99. total += labels.size(0)
  100. correct += (predicted == labels).sum().item()
  101.  
  102. accuracy = 100 * correct / total
  103. print(f'\nAccuracy of the HyperEfficientCNN model on the {total} test images: {accuracy:.2f} %')
Advertisement
Add Comment
Please, Sign In to add comment