import torch.nn.functional as F
from torch import nn


class Projection(nn.Module):
    def __init__(self, input_dim, output_dim, dropout=0.5):
        super(Projection, self).__init__()
        self.linear_1 = nn.Linear(input_dim, output_dim, bias=False)
        self.linear_2 = nn.Linear(output_dim, output_dim, bias=False)
        self.layer_norm = nn.LayerNorm(output_dim)
        self.dropout = nn.Dropout(dropout)

    def forward(self, x):
        emb1 = self.linear_1(x)
        emb2 = self.dropout(self.linear_2(F.gelu(emb1)))
        return self.layer_norm(emb1 + emb2)


class MLP(nn.Module):
    def __init__(self, units=[512, 512, 512], nonlin=nn.ReLU(), dropout=0.1):
        super(MLP, self).__init__()
        self.nonlin = nonlin
        self.dropout = dropout

        sequence = []
        for u0, u1 in zip(units[:-1], units[1:]):
            sequence.append(nn.Linear(u0, u1))
            sequence.append(self.nonlin)
            sequence.append(nn.Dropout(self.dropout))
        sequence = sequence[:-2]

        self.sequential = nn.Sequential(*sequence)

    def forward(self, x):
        return self.sequential(x)