(self, hidden_states: torch.Tensor)
| 185 | self.eps = eps |
| 186 | |
| 187 | def forward(self, hidden_states: torch.Tensor): |
| 188 | input_dtype = hidden_states.dtype |
| 189 | variance = hidden_states.to(torch.float32).pow(2).mean(-1, keepdim=True) |
| 190 | hidden_states = hidden_states * torch.rsqrt(variance + self.eps) |
| 191 | |
| 192 | return (self.weight * hidden_states).to(input_dtype) |
| 193 | |
| 194 | |
| 195 | class CoreAttention(torch.nn.Module): |
nothing calls this directly
no outgoing calls
no test coverage detected