lenet.py 1.5 KB

12345678910111213141516171819202122232425262728293031323334353637383940414243444546474849505152
  1. import torch
  2. from torch import nn
  3. class LeNet5(nn.Module):
  4. def __init__(self,dropout_rate=0.5):
  5. super().__init__()
  6. #特征提取器 Conv → ReLU → Pool 反复出现
  7. #卷积提取局部模式,ReLU 引入非线性,池化降低空间分辨率并扩大有效视野。
  8. self.conv1 = nn.Sequential(
  9. nn.Conv2d(in_channels=3, out_channels=6, kernel_size=5, stride=1, padding=0),
  10. nn.ReLU(),
  11. nn.MaxPool2d(kernel_size=2, stride=2, padding=0)
  12. )
  13. self.conv2 = nn.Sequential(
  14. nn.Conv2d(in_channels=6, out_channels=16, kernel_size=5, stride=1, padding=0),
  15. nn.ReLU(),
  16. nn.MaxPool2d(kernel_size=2, stride=2, padding=0)
  17. )
  18. self.flatten = nn.Flatten()
  19. self.fc1 = nn.Sequential(
  20. nn.Linear(400, 120),
  21. nn.ReLU(),
  22. nn.Dropout(p=dropout_rate)
  23. )
  24. self.fc2 = nn.Sequential(
  25. nn.Linear(120, 84),
  26. nn.ReLU(),
  27. nn.Dropout(p=dropout_rate)
  28. )
  29. self.fc3 = nn.Linear(84, 10)
  30. def forward(self, x):
  31. out = self.conv1(x)
  32. out = self.conv2(out)
  33. out = self.flatten(out)
  34. out = self.fc1(out)
  35. out = self.fc2(out)
  36. logits = self.fc3(out)
  37. return logits
  38. if __name__ == '__main__':
  39. net = LeNet5()
  40. data = torch.randn(1, 3, 32, 32)
  41. output = net(data)
  42. print(output)