| 145 | |
| 146 | |
| 147 | class T5RelativeEmbedding(nn.Module): |
| 148 | |
| 149 | def __init__(self, num_buckets, num_heads, bidirectional, max_dist=128): |
| 150 | super(T5RelativeEmbedding, self).__init__() |
| 151 | self.num_buckets = num_buckets |
| 152 | self.num_heads = num_heads |
| 153 | self.bidirectional = bidirectional |
| 154 | self.max_dist = max_dist |
| 155 | |
| 156 | # layers |
| 157 | self.embedding = nn.Embedding(num_buckets, num_heads) |
| 158 | |
| 159 | def forward(self, lq, lk): |
| 160 | device = self.embedding.weight.device |
| 161 | # rel_pos = torch.arange(lk).unsqueeze(0).to(device) - \ |
| 162 | # torch.arange(lq).unsqueeze(1).to(device) |
| 163 | rel_pos = torch.arange(lk, device=device).unsqueeze(0) - \ |
| 164 | torch.arange(lq, device=device).unsqueeze(1) |
| 165 | rel_pos = self._relative_position_bucket(rel_pos) |
| 166 | rel_pos_embeds = self.embedding(rel_pos) |
| 167 | rel_pos_embeds = rel_pos_embeds.permute(2, 0, 1).unsqueeze( |
| 168 | 0) # [1, N, Lq, Lk] |
| 169 | return rel_pos_embeds.contiguous() |
| 170 | |
| 171 | def _relative_position_bucket(self, rel_pos): |
| 172 | # preprocess |
| 173 | if self.bidirectional: |
| 174 | num_buckets = self.num_buckets // 2 |
| 175 | rel_buckets = (rel_pos > 0).long() * num_buckets |
| 176 | rel_pos = torch.abs(rel_pos) |
| 177 | else: |
| 178 | num_buckets = self.num_buckets |
| 179 | rel_buckets = 0 |
| 180 | rel_pos = -torch.min(rel_pos, torch.zeros_like(rel_pos)) |
| 181 | |
| 182 | # embeddings for small and large positions |
| 183 | max_exact = num_buckets // 2 |
| 184 | rel_pos_large = max_exact + (torch.log(rel_pos.float() / max_exact) / |
| 185 | math.log(self.max_dist / max_exact) * |
| 186 | (num_buckets - max_exact)).long() |
| 187 | rel_pos_large = torch.min( |
| 188 | rel_pos_large, torch.full_like(rel_pos_large, num_buckets - 1)) |
| 189 | rel_buckets += torch.where(rel_pos < max_exact, rel_pos, rel_pos_large) |
| 190 | return rel_buckets |
| 191 | |
| 192 | def init_weights(m): |
| 193 | if isinstance(m, T5LayerNorm): |