net.py 1.0 KB

1234567891011121314151617181920212223242526272829
  1. from torch import nn
  2. import torch
  3. class FullyConnectedNet(nn.Module):
  4. def __init__(self):
  5. super().__init__()
  6. # nn.Sequential 会按顺序执行每一层。
  7. self.layer = nn.Sequential(
  8. nn.Flatten(), # [batch, 1, 28, 28] -> [batch, 784]
  9. nn.Linear(28 * 28, 512), # 784 个像素点映射到 512 个隐藏特征
  10. nn.ReLU(), # 增加非线性表达能力
  11. nn.Linear(512, 256), # 继续提取更紧凑的隐藏特征
  12. nn.ReLU(),
  13. nn.Linear(256, 128),
  14. nn.ReLU(),
  15. nn.Linear(128, 10), # 输出 10 个数字类别的分数
  16. )
  17. def forward(self, x):
  18. # forward 定义“数据如何从输入流到输出”。
  19. # 输入 x 是一批图片,输出是每张图片对应 10 个数字类别的 logits。
  20. return self.layer(x)
  21. if __name__ == '__main__':
  22. data = torch.randn(1,1,28,28)
  23. net = FullyConnectedNet()
  24. output = net(data)
  25. print(output)