(self, lq, lk)
| 217 | self.embedding = nn.Embedding(num_buckets, num_heads) |
| 218 | |
| 219 | def forward(self, lq, lk): |
| 220 | device = self.embedding.weight.device |
| 221 | # rel_pos = torch.arange(lk).unsqueeze(0).to(device) - \ |
| 222 | # torch.arange(lq).unsqueeze(1).to(device) |
| 223 | if torch.device(type="meta") != device: |
| 224 | rel_pos = torch.arange(lk, device=device).unsqueeze(0) - \ |
| 225 | torch.arange(lq, device=device).unsqueeze(1) |
| 226 | else: |
| 227 | rel_pos = torch.arange(lk).unsqueeze(0) - \ |
| 228 | torch.arange(lq).unsqueeze(1) |
| 229 | rel_pos = self._relative_position_bucket(rel_pos) |
| 230 | rel_pos_embeds = self.embedding(rel_pos) |
| 231 | rel_pos_embeds = rel_pos_embeds.permute(2, 0, 1).unsqueeze( |
| 232 | 0) # [1, N, Lq, Lk] |
| 233 | return rel_pos_embeds.contiguous() |
| 234 | |
| 235 | def _relative_position_bucket(self, rel_pos): |
| 236 | # preprocess |
nothing calls this directly
no test coverage detected