Drops labels to enable classifier-free guidance.
(self, labels, force_drop_ids=None)
| 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 |