(self, hidden_states: torch.Tensor)
| 161 | self.eps = eps |
| 162 | |
| 163 | def forward(self, hidden_states: torch.Tensor): |
| 164 | input_dtype = hidden_states.dtype |
| 165 | variance = hidden_states.to(torch.float32).pow(2).mean(-1, keepdim=True) |
| 166 | hidden_states = hidden_states * torch.rsqrt(variance + self.eps) |
| 167 | |
| 168 | return (self.weight * hidden_states).to(input_dtype) |
| 169 | |
| 170 | |
| 171 | class CoreAttention(torch.nn.Module): |
nothing calls this directly
no outgoing calls
no test coverage detected