(self, channels: int, in_channels: int = 3)
| 10 | Embeds spatial positions into vector representations. |
| 11 | """ |
| 12 | def __init__(self, channels: int, in_channels: int = 3): |
| 13 | super().__init__() |
| 14 | self.channels = channels |
| 15 | self.in_channels = in_channels |
| 16 | self.freq_dim = channels // in_channels // 2 |
| 17 | self.freqs = torch.arange(self.freq_dim, dtype=torch.float32) / self.freq_dim |
| 18 | self.freqs = 1.0 / (10000 ** self.freqs) |
| 19 | |
| 20 | def _sin_cos_embedding(self, x: torch.Tensor) -> torch.Tensor: |
| 21 | """ |