Embeds class labels into vector representations. Also handles label dropout for classifier-free guidance.
| 276 | |
| 277 | |
| 278 | class LabelEmbedder(nn.Module): |
| 279 | """ |
| 280 | Embeds class labels into vector representations. Also handles label dropout for classifier-free guidance. |
| 281 | """ |
| 282 | |
| 283 | def __init__(self, num_classes, hidden_size, dropout_prob): |
| 284 | super().__init__() |
| 285 | use_cfg_embedding = dropout_prob > 0 |
| 286 | self.embedding_table = nn.Embedding( |
| 287 | num_classes + use_cfg_embedding, hidden_size |
| 288 | ) |
| 289 | self.num_classes = num_classes |
| 290 | self.dropout_prob = dropout_prob |
| 291 | |
| 292 | def token_drop(self, labels, force_drop_ids=None): |
| 293 | """ |
| 294 | Drops labels to enable classifier-free guidance. |
| 295 | """ |
| 296 | if force_drop_ids is None: |
| 297 | drop_ids = ( |
| 298 | torch.rand(labels.shape[0], device=labels.device) < self.dropout_prob |
| 299 | ) |
| 300 | else: |
| 301 | drop_ids = force_drop_ids == 1 |
| 302 | labels = torch.where(drop_ids, self.num_classes, labels) |
| 303 | return labels |
| 304 | |
| 305 | def forward(self, labels, train, force_drop_ids=None): |
| 306 | use_dropout = self.dropout_prob > 0 |
| 307 | if (train and use_dropout) or (force_drop_ids is not None): |
| 308 | labels = self.token_drop(labels, force_drop_ids) |
| 309 | embeddings = self.embedding_table(labels) |
| 310 | return embeddings |
| 311 | |
| 312 | |
| 313 | class FinalLayer(nn.Module): |