import math
import torch
from torch import nn
def gelu_exact(x):
"""Exact GELU, applied elementwise to a tensor of any shape."""
return 0.5 * x * (1.0 + torch.erf(x / math.sqrt(2.0)))
class FFN(nn.Module):
def __init__(self, d_model, d_ff):
super(FFN, self).__init__()
self.w_1 = nn.Linear(d_model, d_ff)
self.w_2 = nn.Linear(d_ff, d_model)
def forward(self, x):
hidden = self.w_1(x)
activated = gelu_exact(hidden)
return self.w_2(activated)
def gelu_tanh(x):
scale = math.sqrt(2.0 / math.pi)
return 0.5 * x * (
1.0 + torch.tanh(
scale * (x + 0.044715 * x.pow(3))
)
)