lenet.py 2.0 KB

12345678910111213141516171819202122232425262728293031323334353637383940414243444546474849505152535455565758596061626364656667
  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,
  10. out_channels=6,
  11. kernel_size=5,
  12. stride=1,
  13. padding=0,
  14. bias=False),
  15. nn.BatchNorm2d(6), # 在激活函数之前使用BN
  16. nn.ReLU(),
  17. nn.MaxPool2d(kernel_size=2, stride=2, padding=0)
  18. )
  19. self.conv2 = nn.Sequential(
  20. nn.Conv2d(in_channels=6,
  21. out_channels=16,
  22. kernel_size=5,
  23. stride=1,
  24. padding=0,
  25. bias=False),
  26. nn.BatchNorm2d(16), # 在激活函数之前使用BN
  27. nn.ReLU(),
  28. nn.MaxPool2d(kernel_size=2, stride=2, padding=0)
  29. )
  30. self.flatten = nn.Flatten()
  31. self.fc1 = nn.Sequential(
  32. nn.Linear(400, 120,bias=False),
  33. nn.BatchNorm1d(120), # 在激活函数之前使用BN
  34. nn.ReLU(),
  35. nn.Dropout(p=dropout_rate)
  36. )
  37. self.fc2 = nn.Sequential(
  38. nn.Linear(120, 84,bias=False),
  39. nn.BatchNorm1d(84), # 在激活函数之前使用BN
  40. nn.ReLU(),
  41. nn.Dropout(p=dropout_rate)
  42. )
  43. self.fc3 = nn.Linear(84, 10)
  44. def forward(self, x):
  45. out = self.conv1(x)
  46. out = self.conv2(out)
  47. out = self.flatten(out)
  48. out = self.fc1(out)
  49. out = self.fc2(out)
  50. logits = self.fc3(out)
  51. return logits
  52. if __name__ == '__main__':
  53. net = LeNet5(dropout_rate=0.5)
  54. data = torch.randn(1, 3, 32, 32)
  55. output = net(data)
  56. print(output)
  57. print(output.shape)