MCPcopy Create free account
hub / github.com/DFin/Neural-Network-Visualisation / SmallMLP

Class SmallMLP

training/mlp_train.py:44–59  ·  view source on GitHub ↗

Simple fully connected network for MNIST digits.

Source from the content-addressed store, hash-verified

42
43
44class SmallMLP(nn.Module):
45 """Simple fully connected network for MNIST digits."""
46
47 def __init__(self, input_dim: int, hidden_dims: Sequence[int], num_classes: int = 10):
48 super().__init__()
49 dims = [input_dim, *hidden_dims, num_classes]
50 layers: list[nn.Module] = []
51 for idx in range(len(dims) - 1):
52 layers.append(nn.Linear(dims[idx], dims[idx + 1]))
53 if idx < len(dims) - 2:
54 layers.append(nn.ReLU())
55 self.net = nn.Sequential(*layers)
56
57 def forward(self, x: torch.Tensor) -> torch.Tensor: # type: ignore[override]
58 x = x.view(x.size(0), -1)
59 return self.net(x)
60
61
62@dataclass

Callers 1

mainFunction · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected