def forward(self, x): layer = nn.Linear(10, 64) # new layer each call return self.head(torch.relu(layer(x)))