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))
        )
    )