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)