(self, batch, key=None, disable_dropout=False)
| 31 | self.ucg_rate = ucg_rate |
| 32 | |
| 33 | def forward(self, batch, key=None, disable_dropout=False): |
| 34 | if key is None: |
| 35 | key = self.key |
| 36 | # this is for use in crossattn |
| 37 | c = batch[key][:, None] |
| 38 | if self.ucg_rate > 0. and not disable_dropout: |
| 39 | mask = 1. - torch.bernoulli(torch.ones_like(c) * self.ucg_rate) |
| 40 | c = mask * c + (1-mask) * torch.ones_like(c)*(self.n_classes-1) |
| 41 | c = c.long() |
| 42 | c = self.embedding(c) |
| 43 | return c |
| 44 | |
| 45 | def get_unconditional_conditioning(self, bs, device="cuda"): |
| 46 | uc_class = self.n_classes - 1 # 1000 classes --> 0 ... 999, one extra class for ucg (class 1000) |
nothing calls this directly
no outgoing calls
no test coverage detected