hikmatfarhat/MNIST_Classifier
08
1import torch.nn as nn2 3class Net(nn.Module):4 def __init__(self,input_size,hidden_size1,hidden_size2,output_size):5 super(Net, self).__init__()6 self.layer1=nn.Linear(input_size,hidden_size1)7 self.layer2=nn.Linear(hidden_size1,hidden_size2)8 self.layer3=nn.Linear(hidden_size2,output_size)9 self.relu=nn.ReLU()10 def forward(self,x):11 x=x.flatten(start_dim=1)12 x=self.layer1(x)13 x=self.relu(x)14 x=self.layer2(x)15 x=self.relu(x)16 x=self.layer3(x)17 18 return x