| 19 | |
| 20 | @register_norm_module |
| 21 | class LayerNorm(nn.Module): |
| 22 | def __init__(self, hidden_size, eps=1e-12): |
| 23 | """Construct a layernorm module in the TF style (epsilon inside the square root). |
| 24 | """ |
| 25 | super(LayerNorm, self).__init__() |
| 26 | self.weight = nn.Parameter(torch.ones(hidden_size)) |
| 27 | self.bias = nn.Parameter(torch.zeros(hidden_size)) |
| 28 | self.variance_epsilon = eps |
| 29 | |
| 30 | def forward(self, x): |
| 31 | pdtype = x.dtype |
| 32 | x = x.float() |
| 33 | u = x.mean(-1, keepdim=True) |
| 34 | s = (x - u).pow(2).mean(-1, keepdim=True) |
| 35 | x = (x - u) / torch.sqrt(s + self.variance_epsilon) |
| 36 | return self.weight * x.to(pdtype) + self.bias |
| 37 | |
| 38 | |
| 39 | class QuickGELU(nn.Module): |