(self, hidden_states)
| 35 | self.variance_epsilon = eps |
| 36 | |
| 37 | def forward(self, hidden_states): |
| 38 | input_dtype = hidden_states.dtype |
| 39 | variance = hidden_states.to(torch.float32).pow(2).mean(-1, keepdim=True) |
| 40 | hidden_states = hidden_states * torch.rsqrt(variance + self.variance_epsilon) |
| 41 | |
| 42 | return (self.weight * hidden_states).to(input_dtype) |