lenet.py 1.6 KB

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