| 155 | |
| 156 | |
| 157 | class RMSNorm(torch.nn.Module): |
| 158 | def __init__(self, normalized_shape, eps=1e-5, device=None, dtype=None, **kwargs): |
| 159 | super().__init__() |
| 160 | self.weight = torch.nn.Parameter(torch.empty(normalized_shape, device=device, dtype=dtype)) |
| 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