(self, x)
| 38 | self.gamma = nn.Parameter(torch.ones(heads, 1, dim) / self.scale) |
| 39 | |
| 40 | def forward(self, x): |
| 41 | normed = F.normalize(x, dim=-1) |
| 42 | return normed * self.scale * self.gamma |
| 43 | |
| 44 | |
| 45 | # classes |
nothing calls this directly
no outgoing calls
no test coverage detected