(self, vocab_ids)
| 167 | fan_mode=fan_mode) |
| 168 | |
| 169 | def forward(self, vocab_ids): |
| 170 | seq_length = vocab_ids.size(1) |
| 171 | actual_length = vocab_ids.size(1) - self.radius * 2 |
| 172 | trim_vocab_id = vocab_ids[:, self.radius:seq_length - self.radius] |
| 173 | slice_vocabs = \ |
| 174 | [vocab_ids[:, i:i + self.region_size] for i in |
| 175 | range(actual_length)] |
| 176 | slice_vocabs = torch.cat(slice_vocabs, 1) |
| 177 | slice_vocabs = \ |
| 178 | slice_vocabs.view(-1, actual_length, self.region_size) |
| 179 | |
| 180 | if self.region_embedding_type == RegionEmbeddingType.WC: |
| 181 | vocab_embedding = self.embedding(slice_vocabs) |
| 182 | context_embedding = self.context_embedding(trim_vocab_id) |
| 183 | context_embedding = context_embedding.view( |
| 184 | -1, actual_length, self.region_size, self.embedding_dim) |
| 185 | region_embedding = vocab_embedding * context_embedding |
| 186 | region_embedding, _ = region_embedding.max(2) |
| 187 | elif self.region_embedding_type == RegionEmbeddingType.CW: |
| 188 | vocab_embedding = self.embedding(trim_vocab_id).unsqueeze(2) |
| 189 | context_embedding = self.context_embedding(slice_vocabs) |
| 190 | size = context_embedding.size() |
| 191 | context_embedding = context_embedding.view( |
| 192 | size[0], size[1], size[2], self.region_size, self.embedding_dim) |
| 193 | mask = torch.ones( |
| 194 | [self.region_size, self.region_size, self.embedding_dim]) |
| 195 | |
| 196 | for i in range(self.region_size): |
| 197 | mask[i][self.region_size - i - 1] = 0. |
| 198 | neg_mask = mask * -65500.0 |
| 199 | mask = mask.le(0).float() |
| 200 | mask = mask.unsqueeze(0).unsqueeze(0) |
| 201 | context_embedding = context_embedding * mask |
| 202 | context_embedding = context_embedding + neg_mask |
| 203 | context_embedding, _ = context_embedding.max(3) |
| 204 | region_embedding = vocab_embedding * context_embedding |
| 205 | region_embedding, _ = region_embedding.max(2) |
| 206 | else: |
| 207 | raise TypeError( |
| 208 | "Unsupported region embedding type: %s." % |
| 209 | self.region_embedding_type) |
| 210 | |
| 211 | return region_embedding |
| 212 | |
| 213 | class PositionEmbedding(torch.nn.Module): |
| 214 | ''' Reference: attention is all you need ''' |
nothing calls this directly
no outgoing calls
no test coverage detected