r"""Embeds class labels into vector representations. Also handles label dropout for classifier-free guidance.
| 85 | |
| 86 | |
| 87 | class ParallelLabelEmbedder(nn.Module): |
| 88 | r"""Embeds class labels into vector representations. Also handles label |
| 89 | dropout for classifier-free guidance. |
| 90 | """ |
| 91 | def __init__(self, num_classes, hidden_size, dropout_prob): |
| 92 | super().__init__() |
| 93 | use_cfg_embedding = int(dropout_prob > 0) |
| 94 | self.embedding_table = ParallelEmbedding( |
| 95 | num_classes + use_cfg_embedding, hidden_size, |
| 96 | init_method=functools.partial(nn.init.normal_, std=0.02), |
| 97 | ) |
| 98 | self.num_classes = num_classes |
| 99 | self.dropout_prob = dropout_prob |
| 100 | |
| 101 | def token_drop(self, labels, force_drop_ids=None): |
| 102 | """ |
| 103 | Drops labels to enable classifier-free guidance. |
| 104 | """ |
| 105 | if force_drop_ids is None: |
| 106 | drop_ids = torch.rand( |
| 107 | labels.shape[0], device=labels.device |
| 108 | ) < self.dropout_prob |
| 109 | drop_ids = drop_ids.cuda() |
| 110 | dist.broadcast( |
| 111 | drop_ids, |
| 112 | fs_init.get_model_parallel_src_rank(), |
| 113 | fs_init.get_model_parallel_group(), |
| 114 | ) |
| 115 | drop_ids = drop_ids.to(labels.device) |
| 116 | else: |
| 117 | drop_ids = force_drop_ids == 1 |
| 118 | labels = torch.where(drop_ids, self.num_classes, labels) |
| 119 | return labels |
| 120 | |
| 121 | def forward(self, labels, train, force_drop_ids=None): |
| 122 | use_dropout = self.dropout_prob > 0 |
| 123 | if (train and use_dropout) or (force_drop_ids is not None): |
| 124 | labels = self.token_drop(labels, force_drop_ids) |
| 125 | embeddings = self.embedding_table(labels) |
| 126 | return embeddings |
| 127 | |
| 128 | |
| 129 | ############################################################################# |