tarekfer8's picture
Upload folder using huggingface_hub
bf8df4f verified
Raw
History Blame Contribute Delete
1.44 kB
import torch.nn as nn
import torch
class MLP(nn.Module):
def __init__(self, input_size, output_size, n_neurons, dropout_rates):
super(MLP, self).__init__()
self.input_size = input_size
self.output_size = output_size
self.n_neurons = n_neurons
self.dropout_rates = dropout_rates
self.fc = nn.ModuleList()
self.dropout = nn.ModuleList()
self.fc.append(nn.Linear(input_size, n_neurons[0]))
self.dropout.append(nn.Dropout(dropout_rates[0]))
for i in range(1, len(n_neurons)):
self.fc.append(nn.Linear(n_neurons[i-1], n_neurons[i]))
self.dropout.append(nn.Dropout(dropout_rates[i]))
self.fc_out = nn.Linear(n_neurons[-1], output_size)
def forward(self, x, apply_activation= False):
x = nn.functional.relu(self.fc[0](x))
x = self.dropout[0](x)
for i in range(1, len(self.fc)):
x = nn.functional.relu(self.fc[i](x))
x = self.dropout[i](x)
x = self.fc_out(x)
# Apply activation only if requested (for inference)
if apply_activation:
x = torch.sigmoid(x)
return x
def predict(self, x, apply_activation=True):
"""Convenience method for inference with activations applied"""
self.eval()
with torch.no_grad():
return self.forward(x, apply_activation=apply_activation)