| 219 | |
| 220 | |
| 221 | class T5RelativeEmbedding(nn.Module): |
| 222 | |
| 223 | def __init__(self, num_buckets, num_heads, bidirectional, max_dist=128): |
| 224 | super(T5RelativeEmbedding, self).__init__() |
| 225 | self.num_buckets = num_buckets |
| 226 | self.num_heads = num_heads |
| 227 | self.bidirectional = bidirectional |
| 228 | self.max_dist = max_dist |
| 229 | |
| 230 | # layers |
| 231 | self.embedding = nn.Embedding(num_buckets, num_heads) |
| 232 | |
| 233 | def forward(self, lq, lk): |
| 234 | device = self.embedding.weight.device |
| 235 | # rel_pos = torch.arange(lk).unsqueeze(0).to(device) - \ |
| 236 | # torch.arange(lq).unsqueeze(1).to(device) |
| 237 | rel_pos = torch.arange(lk, device=device).unsqueeze(0) - \ |
| 238 | torch.arange(lq, device=device).unsqueeze(1) |
| 239 | rel_pos = self._relative_position_bucket(rel_pos) |
| 240 | rel_pos_embeds = self.embedding(rel_pos) |
| 241 | rel_pos_embeds = rel_pos_embeds.permute(2, 0, 1).unsqueeze( |
| 242 | 0) # [1, N, Lq, Lk] |
| 243 | return rel_pos_embeds.contiguous() |
| 244 | |
| 245 | def _relative_position_bucket(self, rel_pos): |
| 246 | # preprocess |
| 247 | if self.bidirectional: |
| 248 | num_buckets = self.num_buckets // 2 |
| 249 | rel_buckets = (rel_pos > 0).long() * num_buckets |
| 250 | rel_pos = torch.abs(rel_pos) |
| 251 | else: |
| 252 | num_buckets = self.num_buckets |
| 253 | rel_buckets = 0 |
| 254 | rel_pos = -torch.min(rel_pos, torch.zeros_like(rel_pos)) |
| 255 | |
| 256 | # embeddings for small and large positions |
| 257 | max_exact = num_buckets // 2 |
| 258 | rel_pos_large = max_exact + (torch.log(rel_pos.float() / max_exact) / |
| 259 | math.log(self.max_dist / max_exact) * |
| 260 | (num_buckets - max_exact)).long() |
| 261 | rel_pos_large = torch.min( |
| 262 | rel_pos_large, torch.full_like(rel_pos_large, num_buckets - 1)) |
| 263 | rel_buckets += torch.where(rel_pos < max_exact, rel_pos, rel_pos_large) |
| 264 | return rel_buckets |
| 265 | |
| 266 | |
| 267 | class T5Encoder(nn.Module): |